1475 lines
46 KiB
Go
1475 lines
46 KiB
Go
package backend
|
|
|
|
// KiroBackend implements Backend for the Kiro streaming API.
|
|
//
|
|
// COVERAGE: Intentionally untested. Reverse-engineered Kiro streaming API
|
|
// client requiring a live session to test.
|
|
//
|
|
//
|
|
// Auth is configured via the apiKey parameter to NewKiro:
|
|
// - 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"
|
|
"crypto/rand"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"runtime"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"ollie/util"
|
|
)
|
|
|
|
// ── Public constructor ────────────────────────────────────────────────────────
|
|
|
|
type KiroBackend struct {
|
|
baseBackend
|
|
endpoint string // overrides auth-derived endpoint if non-empty
|
|
extraHeaders map[string]string
|
|
authSource kiroAuthSource
|
|
authInitErr error
|
|
httpClient *http.Client
|
|
|
|
// Models cache at the backend level (shared across all sessions).
|
|
modelsMu sync.Mutex
|
|
modelsCache []string
|
|
modelsCacheAt time.Time
|
|
}
|
|
|
|
// NewKiro returns a Kiro backend. See package doc for apiKey semantics.
|
|
func NewKiro(apiKey string) (*KiroBackend, error) {
|
|
if strings.TrimSpace(apiKey) == "" {
|
|
apiKey = kiroDefaultSQLiteKey()
|
|
}
|
|
authSource, err := newKiroAuthSource(apiKey)
|
|
b := &KiroBackend{
|
|
baseBackend: baseBackend{name: "kiro"},
|
|
authSource: authSource,
|
|
authInitErr: err,
|
|
httpClient: sharedClient("kiro:https://codewhisperer.us-east-1.amazonaws.com"),
|
|
}
|
|
b.model = b.DefaultModel()
|
|
return b, nil
|
|
}
|
|
|
|
func (b *KiroBackend) DefaultModel() string { return "auto" }
|
|
|
|
func (b *KiroBackend) Models(ctx context.Context) []string {
|
|
b.modelsMu.Lock()
|
|
if b.modelsCache != nil && time.Since(b.modelsCacheAt) < 24*time.Hour {
|
|
result := b.modelsCache
|
|
b.modelsMu.Unlock()
|
|
return result
|
|
}
|
|
b.modelsMu.Unlock()
|
|
|
|
resp := b.fetchModels(ctx)
|
|
if resp == nil {
|
|
return nil
|
|
}
|
|
ids := make([]string, len(resp.Models))
|
|
for i, m := range resp.Models {
|
|
ids[i] = m.ModelID
|
|
}
|
|
|
|
b.modelsMu.Lock()
|
|
b.modelsCache = ids
|
|
b.modelsCacheAt = time.Now()
|
|
b.modelsMu.Unlock()
|
|
return ids
|
|
}
|
|
|
|
func (b *KiroBackend) fetchModels(ctx context.Context) *kiroListModelsResponse {
|
|
if b.authInitErr != nil {
|
|
return nil
|
|
}
|
|
token, err := b.authSource.AccessToken(ctx)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
endpoint, err := b.resolveEndpoint(ctx)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
profileARN, _ := b.authSource.ProfileARN(ctx)
|
|
client := newKiroAPIClient(endpoint, token, b.extraHeaders, b.httpClient)
|
|
resp, err := client.ListModels(ctx, profileARN)
|
|
if err != nil {
|
|
return nil
|
|
}
|
|
return resp
|
|
}
|
|
|
|
func (b *KiroBackend) ContextLength(ctx context.Context) int {
|
|
if b.ctxLen > 0 && b.ctxModel == b.model {
|
|
return b.ctxLen
|
|
}
|
|
resp := b.fetchModels(ctx)
|
|
if resp == nil {
|
|
return 0
|
|
}
|
|
for _, m := range resp.Models {
|
|
if m.TokenLimits != nil && (b.model == "auto" || m.ModelID == b.model) {
|
|
b.ctxLen = m.TokenLimits.MaxInputTokens
|
|
b.ctxModel = b.model
|
|
return b.ctxLen
|
|
}
|
|
}
|
|
if resp.DefaultModel != nil && resp.DefaultModel.TokenLimits != nil {
|
|
b.ctxLen = resp.DefaultModel.TokenLimits.MaxInputTokens
|
|
b.ctxModel = b.model
|
|
return b.ctxLen
|
|
}
|
|
return 0
|
|
}
|
|
|
|
// ── Backend interface ─────────────────────────────────────────────────────────
|
|
|
|
func (b *KiroBackend) ChatStream(
|
|
ctx context.Context,
|
|
messages []Message,
|
|
tools []Tool,
|
|
params GenerationParams,
|
|
) (<-chan StreamEvent, error) {
|
|
model := b.model
|
|
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, params)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
ch := make(chan StreamEvent, 8)
|
|
go func() {
|
|
defer close(ch)
|
|
b.runStream(ctx, req, ch)
|
|
}()
|
|
return ch, nil
|
|
}
|
|
|
|
// kiroMaxThrottleRetries is the maximum number of retries on throttling errors.
|
|
const kiroMaxThrottleRetries = 5
|
|
|
|
// kiroThrottleBaseDelay is the initial backoff delay for throttling retries.
|
|
// Tests can override this.
|
|
var kiroThrottleBaseDelay = 5 * time.Second
|
|
|
|
func (b *KiroBackend) runStream(ctx context.Context, req *kiroGenerateRequest, ch chan<- StreamEvent) {
|
|
var throttleAttempts int
|
|
for {
|
|
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
|
|
}
|
|
},
|
|
OnReasoningDelta: func(event kiroReasoningContentEvent) {
|
|
if stopped || (event.Text == "" && event.RedactedContent == "" && event.Signature == "") {
|
|
return
|
|
}
|
|
select {
|
|
case ch <- StreamEvent{Reasoning: event.Text}:
|
|
case <-ctx.Done():
|
|
stopped = true
|
|
}
|
|
},
|
|
})
|
|
|
|
if streamErr != nil {
|
|
var apiErr *kiroAPIError
|
|
if errors.As(streamErr, &apiErr) && (apiErr.isThrottled() || apiErr.isTransient()) && throttleAttempts < kiroMaxThrottleRetries {
|
|
throttleAttempts++
|
|
var wait time.Duration
|
|
if apiErr.isThrottled() {
|
|
wait = kiroThrottleBaseDelay << (throttleAttempts - 1)
|
|
} else {
|
|
wait = time.Duration(throttleAttempts) * time.Second
|
|
}
|
|
select {
|
|
case <-ctx.Done():
|
|
ch <- StreamEvent{Done: true, StopReason: ctx.Err().Error()}
|
|
return
|
|
case <-time.After(wait):
|
|
}
|
|
continue
|
|
}
|
|
if throttleAttempts == 0 && b.shouldRefresh(streamErr) {
|
|
if refreshErr := b.authSource.Refresh(ctx); refreshErr != nil {
|
|
ch <- StreamEvent{Done: true, StopReason: fmt.Sprintf("token refresh: %v", refreshErr)}
|
|
return
|
|
}
|
|
throttleAttempts++ // use throttleAttempts as general attempt counter for auth retry
|
|
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 *KiroBackend) resolveEndpoint(ctx context.Context) (string, error) {
|
|
if ep := strings.TrimSpace(b.endpoint); ep != "" {
|
|
return ep, nil
|
|
}
|
|
return b.authSource.DefaultEndpoint(ctx)
|
|
}
|
|
|
|
func (b *KiroBackend) 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, params GenerationParams) (*kiroGenerateRequest, error) {
|
|
model = kiroSanitizeModel(model)
|
|
if len(messages) == 0 {
|
|
return nil, fmt.Errorf("at least one message is required")
|
|
}
|
|
|
|
// Bedrock requires toolConfig when toolUse/toolResult blocks are present.
|
|
// If no tools are defined, strip tool interactions from history to avoid
|
|
// "The toolConfig field must be defined when using toolUse and toolResult
|
|
// content blocks" errors.
|
|
if len(tools) == 0 {
|
|
messages = kiroStripToolBlocks(messages)
|
|
}
|
|
|
|
// Apply counter-thinking tags to user messages before assistant responses.
|
|
messages = kiroApplyCounterThinkingTags(messages)
|
|
|
|
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, params.CWD)}
|
|
|
|
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, images := encodeKiroToolResults(messages[cutoff+1:])
|
|
current = kiroChatMessage{UserInputMessage: emptyKiroUserMessage(model, encodedTools, toolResults, images, params.CWD)}
|
|
|
|
default:
|
|
return nil, fmt.Errorf("last message role %q is not supported (expected user, system, or tool)", lastRole)
|
|
}
|
|
|
|
convID := util.NewUUID()
|
|
contID := util.NewUUID()
|
|
|
|
req := &kiroGenerateRequest{
|
|
ProfileARN: profileARN,
|
|
ConversationState: kiroConversationState{
|
|
ConversationID: convID,
|
|
History: history,
|
|
CurrentMessage: current,
|
|
ChatTriggerType: kiroTriggerType,
|
|
AgentContinuationID: contID,
|
|
AgentTaskType: kiroAgentTaskType,
|
|
},
|
|
}
|
|
|
|
// Set additionalModelRequestFields for reasoning/thinking support.
|
|
if fields := kiroModelRequestFieldsFromParams(params); fields != nil {
|
|
req.AdditionalModelRequestFields = fields
|
|
}
|
|
|
|
return req, 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, imgs := encodeKiroToolResults(messages[start : i+1])
|
|
out = append(out, kiroChatMessage{UserInputMessage: emptyKiroUserMessage("", nil, results, imgs, "")})
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func encodeKiroUserMessage(m Message, modelID string, tools []kiroTool, toolResults []kiroToolResult, cwd string) *kiroUserInputMessage {
|
|
msg := &kiroUserInputMessage{
|
|
Content: kiroExtractTextContent(m),
|
|
Origin: kiroDefaultOrigin,
|
|
ModelID: modelID,
|
|
Images: kiroExtractImages(m),
|
|
}
|
|
ctx := &kiroUserInputContext{}
|
|
if len(tools) > 0 {
|
|
ctx.Tools = tools
|
|
}
|
|
if len(toolResults) > 0 {
|
|
ctx.ToolResults = toolResults
|
|
}
|
|
if cwd != "" {
|
|
ctx.EnvState = kiroCurrentEnvState(cwd)
|
|
}
|
|
if len(ctx.Tools) > 0 || len(ctx.ToolResults) > 0 || ctx.EnvState != nil {
|
|
msg.UserInputMessageContext = ctx
|
|
}
|
|
return msg
|
|
}
|
|
|
|
func emptyKiroUserMessage(modelID string, tools []kiroTool, toolResults []kiroToolResult, images []kiroImageBlock, cwd string) *kiroUserInputMessage {
|
|
msg := &kiroUserInputMessage{
|
|
Content: "",
|
|
Origin: kiroDefaultOrigin,
|
|
ModelID: modelID,
|
|
Images: images,
|
|
}
|
|
ctx := &kiroUserInputContext{}
|
|
if len(tools) > 0 {
|
|
ctx.Tools = tools
|
|
}
|
|
if len(toolResults) > 0 {
|
|
ctx.ToolResults = toolResults
|
|
}
|
|
if cwd != "" {
|
|
ctx.EnvState = kiroCurrentEnvState(cwd)
|
|
}
|
|
if len(ctx.Tools) > 0 || len(ctx.ToolResults) > 0 || ctx.EnvState != nil {
|
|
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("{}")
|
|
} else {
|
|
// Normalize the JSON by round-tripping through unmarshal/marshal.
|
|
// This deduplicates keys (the model sometimes produces duplicate
|
|
// keys which the upstream API rejects) and ensures well-formed output.
|
|
var parsed any
|
|
if err := json.Unmarshal(input, &parsed); err == nil {
|
|
if normalized, err := json.Marshal(parsed); err == nil {
|
|
input = normalized
|
|
}
|
|
}
|
|
}
|
|
toolUseID := strings.TrimSpace(tc.ID)
|
|
if toolUseID == "" {
|
|
toolUseID = util.NewUUID()
|
|
}
|
|
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, []kiroImageBlock) {
|
|
results := make([]kiroToolResult, 0, len(messages))
|
|
var images []kiroImageBlock
|
|
for _, m := range messages {
|
|
if !strings.EqualFold(m.Role, "tool") {
|
|
continue
|
|
}
|
|
toolUseID := strings.TrimSpace(m.ToolCallID)
|
|
content, imgs := encodeKiroToolResultBlocks(m)
|
|
results = append(results, kiroToolResult{
|
|
ToolUseID: toolUseID,
|
|
Content: content,
|
|
})
|
|
images = append(images, imgs...)
|
|
}
|
|
return results, images
|
|
}
|
|
|
|
func encodeKiroToolResultBlocks(m Message) ([]kiroToolResultContent, []kiroImageBlock) {
|
|
if len(m.ContentBlocks) == 0 {
|
|
return []kiroToolResultContent{encodeKiroToolResultContent(m.Content)}, nil
|
|
}
|
|
out := make([]kiroToolResultContent, 0, len(m.ContentBlocks))
|
|
var images []kiroImageBlock
|
|
for _, cb := range m.ContentBlocks {
|
|
switch cb.Type {
|
|
case "text":
|
|
out = append(out, encodeKiroToolResultContent(cb.Text))
|
|
case "image":
|
|
if cb.ImageSource == nil {
|
|
continue
|
|
}
|
|
format := strings.TrimPrefix(cb.ImageSource.MediaType, "image/")
|
|
images = append(images, kiroImageBlock{
|
|
Format: format,
|
|
Source: kiroImageSource{Bytes: cb.ImageSource.Data},
|
|
})
|
|
}
|
|
}
|
|
if len(out) == 0 {
|
|
out = []kiroToolResultContent{{Text: ""}}
|
|
}
|
|
return out, images
|
|
}
|
|
|
|
func encodeKiroToolResultContent(content string) kiroToolResultContent {
|
|
// Always use the text field. The upstream Amazon Q API rejects bare JSON
|
|
// primitives (numbers, strings, booleans, null) in the "json" field with
|
|
// "ValidationException: Improperly formed request." — it expects only
|
|
// objects or arrays. Using text is always safe; the model sees the same
|
|
// content either way.
|
|
return kiroToolResultContent{Text: strings.TrimSpace(content)}
|
|
}
|
|
|
|
// kiroExtractTextContent returns the text content of a message. When images are
|
|
// present, they are encoded as separate image blocks on the message; this
|
|
// function only returns the text portion.
|
|
func kiroExtractTextContent(m Message) string {
|
|
if len(m.ContentBlocks) == 0 {
|
|
return m.Content
|
|
}
|
|
var sb strings.Builder
|
|
for _, cb := range m.ContentBlocks {
|
|
if cb.Type == "text" {
|
|
sb.WriteString(cb.Text)
|
|
}
|
|
}
|
|
return sb.String()
|
|
}
|
|
|
|
// kiroExtractImages extracts image blocks from a message's ContentBlocks.
|
|
func kiroExtractImages(m Message) []kiroImageBlock {
|
|
if len(m.ContentBlocks) == 0 {
|
|
return nil
|
|
}
|
|
var images []kiroImageBlock
|
|
for _, cb := range m.ContentBlocks {
|
|
if cb.Type == "image" && cb.ImageSource != nil {
|
|
format := strings.TrimPrefix(cb.ImageSource.MediaType, "image/")
|
|
images = append(images, kiroImageBlock{
|
|
Format: format,
|
|
Source: kiroImageSource{Bytes: cb.ImageSource.Data},
|
|
})
|
|
}
|
|
}
|
|
return images
|
|
}
|
|
|
|
func kiroCurrentEnvState(cwd string) *kiroEnvState {
|
|
return &kiroEnvState{
|
|
OperatingSystem: runtime.GOOS,
|
|
CurrentWorkingDirectory: cwd,
|
|
}
|
|
}
|
|
|
|
// ── Tool block stripping ─────────────────────────────────────────────────────
|
|
|
|
// kiroStripToolBlocks removes tool-related content from messages so that
|
|
// the request can be sent without tool definitions. Bedrock rejects requests
|
|
// that contain toolUse/toolResult blocks without a corresponding toolConfig.
|
|
//
|
|
// Strategy:
|
|
// - Assistant messages with ToolCalls: keep only the text content, drop ToolCalls.
|
|
// - Tool-role messages: drop entirely (they are responses to now-removed tool calls).
|
|
// - All other messages: pass through unchanged.
|
|
func kiroStripToolBlocks(messages []Message) []Message {
|
|
// Fast path: if no tool content exists, return as-is.
|
|
hasToolContent := false
|
|
for _, m := range messages {
|
|
if strings.EqualFold(m.Role, "tool") || len(m.ToolCalls) > 0 {
|
|
hasToolContent = true
|
|
break
|
|
}
|
|
}
|
|
if !hasToolContent {
|
|
return messages
|
|
}
|
|
|
|
out := make([]Message, 0, len(messages))
|
|
for _, m := range messages {
|
|
role := strings.ToLower(strings.TrimSpace(m.Role))
|
|
switch {
|
|
case role == "tool":
|
|
// Drop tool result messages entirely.
|
|
continue
|
|
case role == "assistant" && len(m.ToolCalls) > 0:
|
|
// Keep text content, drop tool calls.
|
|
stripped := m
|
|
stripped.ToolCalls = nil
|
|
// If the assistant message was only tool calls with no text, skip it
|
|
// to avoid sending an empty assistant turn.
|
|
if strings.TrimSpace(stripped.Content) == "" && len(stripped.ContentBlocks) == 0 {
|
|
continue
|
|
}
|
|
out = append(out, stripped)
|
|
default:
|
|
out = append(out, m)
|
|
}
|
|
}
|
|
return out
|
|
}
|
|
|
|
// ── Counter-thinking & reasoning support ──────────────────────────────────────
|
|
|
|
const kiroCounterThinkingInstructions = `Be thorough in your work; don't skip steps or cut corners to save time. Think deeply and step by step, do not skip ahead prematurely.`
|
|
const kiroDefaultReasoningEffort = "95"
|
|
|
|
// kiroApplyCounterThinkingTags appends corrective tags to user messages that
|
|
// immediately precede assistant responses, countering server-side injections
|
|
// that degrade reasoning quality.
|
|
func kiroApplyCounterThinkingTags(messages []Message) []Message {
|
|
if len(messages) < 2 {
|
|
return messages
|
|
}
|
|
var mangled []Message
|
|
for i := 1; i < len(messages); i++ {
|
|
if !strings.EqualFold(messages[i].Role, "assistant") {
|
|
continue
|
|
}
|
|
prev := messages[i-1]
|
|
role := strings.ToLower(strings.TrimSpace(prev.Role))
|
|
if role != "user" && role != "system" {
|
|
continue
|
|
}
|
|
if mangled == nil {
|
|
mangled = make([]Message, len(messages))
|
|
copy(mangled, messages)
|
|
}
|
|
suffix := "\n\n<reasoning_effort>" + kiroDefaultReasoningEffort + "</reasoning_effort>\n" +
|
|
"<reasoning_effort>" + kiroDefaultReasoningEffort + "</reasoning_effort>\n" +
|
|
"<implicitInstruction>" + kiroCounterThinkingInstructions + "</implicitInstruction>"
|
|
mangled[i-1].Content = prev.Content + suffix
|
|
}
|
|
if mangled != nil {
|
|
return mangled
|
|
}
|
|
return messages
|
|
}
|
|
|
|
// kiroModelRequestFieldsFromParams builds additionalModelRequestFields from
|
|
// GenerationParams. Returns nil if no reasoning/thinking fields are needed.
|
|
func kiroModelRequestFieldsFromParams(params GenerationParams) *kiroModelRequestFields {
|
|
fields := &kiroModelRequestFields{}
|
|
hasFields := false
|
|
|
|
// Enable summarized thinking by default for reasoning-capable models.
|
|
if params.ThinkingBudget > 0 || params.ReasoningEffort != "" {
|
|
fields.Thinking = &kiroThinkingConfig{
|
|
Type: "adaptive",
|
|
Display: "summarized",
|
|
}
|
|
hasFields = true
|
|
}
|
|
|
|
// Map reasoning effort.
|
|
if effort := kiroMapEffort(params.ReasoningEffort); effort != "" {
|
|
fields.OutputConfig = &kiroOutputConfig{Effort: effort}
|
|
hasFields = true
|
|
} else if params.ThinkingBudget > 0 {
|
|
// Default to high effort when thinking is requested but no specific effort.
|
|
fields.OutputConfig = &kiroOutputConfig{Effort: "high"}
|
|
hasFields = true
|
|
}
|
|
|
|
if !hasFields {
|
|
return nil
|
|
}
|
|
return fields
|
|
}
|
|
|
|
func kiroMapEffort(effort string) string {
|
|
switch strings.ToLower(strings.TrimSpace(effort)) {
|
|
case "max":
|
|
return "max"
|
|
case "high":
|
|
return "high"
|
|
case "medium":
|
|
return "medium"
|
|
case "low", "minimal":
|
|
return "low"
|
|
default:
|
|
return ""
|
|
}
|
|
}
|
|
|
|
// ── 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
|
|
profileARN string
|
|
}
|
|
|
|
func (s *staticKiroAuth) DefaultEndpoint(_ context.Context) (string, error) {
|
|
if s.profileARN != "" {
|
|
return kiroDefaultEndpoint(s.profileARN)
|
|
}
|
|
return "", fmt.Errorf("base URL not configured and no profile ARN available with static auth")
|
|
}
|
|
|
|
func (s *staticKiroAuth) ProfileARN(_ context.Context) (string, error) {
|
|
return s.profileARN, 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
|
|
tokenReadAt time.Time // last time tokens were read from sqlite
|
|
// jitter is computed once at creation time so all checks within this
|
|
// process use the same refresh threshold. Different processes get
|
|
// different jitter values, spreading refresh attempts over time.
|
|
jitter time.Duration
|
|
}
|
|
|
|
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
|
|
}
|
|
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 != "" && !s.tokenExpiringSoon(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()
|
|
|
|
// Bound the refresh operation to avoid hanging indefinitely.
|
|
refreshCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
|
defer cancel()
|
|
|
|
// Always perform a full reload of all sqlite state (including device
|
|
// registration) before refreshing. Another process (e.g. kiro-cli login)
|
|
// may have rotated the client credentials since our initial load.
|
|
state, err := s.fullLoad(refreshCtx)
|
|
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(refreshCtx, s.path)
|
|
if err == nil && currentAccess != "" && currentAccess != state.AccessToken && !s.tokenExpiringSoon(currentExpiry) {
|
|
return nil // already refreshed by another process
|
|
}
|
|
|
|
refreshToken := currentRefresh
|
|
if refreshToken == "" {
|
|
refreshToken = state.RefreshToken
|
|
}
|
|
|
|
if state.IsSocial {
|
|
return s.refreshSocial(refreshCtx, state, refreshToken)
|
|
}
|
|
return s.refreshOIDC(refreshCtx, state, refreshToken, currentAccess)
|
|
}
|
|
|
|
func (s *sqliteKiroAuth) refreshOIDC(ctx context.Context, state kiroSQLiteState, refreshToken, currentAccess string) error {
|
|
oidcEndpoint := fmt.Sprintf("https://oidc.%s.amazonaws.com/token", strings.TrimSpace(state.Region))
|
|
client := newKiroOIDCClient(oidcEndpoint, nil)
|
|
|
|
resp, err := client.RefreshToken(ctx, kiroRefreshTokenRequest{
|
|
ClientID: state.ClientID,
|
|
ClientSecret: state.ClientSecret,
|
|
GrantType: "refresh_token",
|
|
RefreshToken: refreshToken,
|
|
})
|
|
if err != nil {
|
|
// Try SSO cache fallback.
|
|
if fallbackErr := s.refreshFromSSOCache(ctx, currentAccess); fallbackErr == nil {
|
|
return nil
|
|
}
|
|
// Check if another process refreshed while we were trying.
|
|
if s.sqliteHasNewUsableAccessToken(ctx, currentAccess) {
|
|
return 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)
|
|
}
|
|
if err := kiroUpdateSQLiteTokens(ctx, s.path, newAccess, newRefresh, expiresAt); err != nil {
|
|
if s.sqliteHasNewUsableAccessToken(ctx, currentAccess) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *sqliteKiroAuth) refreshSocial(ctx context.Context, state kiroSQLiteState, refreshToken string) error {
|
|
client := newKiroSocialAuthClient(nil)
|
|
resp, err := client.RefreshToken(ctx, 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.RFC3339Nano)
|
|
}
|
|
profileARN := strings.TrimSpace(resp.ProfileARN)
|
|
if profileARN == "" {
|
|
profileARN = state.ProfileARN
|
|
}
|
|
return kiroUpdateSQLiteSocialTokens(ctx, s.path, newAccess, newRefresh, expiresAt, profileARN)
|
|
}
|
|
|
|
func (s *sqliteKiroAuth) CanRefresh() bool {
|
|
state, err := s.load(context.Background())
|
|
return err == nil && s.canRefreshState(state)
|
|
}
|
|
|
|
func (s *sqliteKiroAuth) canRefreshState(state kiroSQLiteState) bool {
|
|
if strings.TrimSpace(state.RefreshToken) == "" {
|
|
return false
|
|
}
|
|
if state.IsSocial {
|
|
return true
|
|
}
|
|
return 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()
|
|
|
|
// Re-read tokens if stale (>60s). Within that window, use cached values
|
|
// to avoid shelling out to sqlite3 on every API call.
|
|
s.mu.Lock()
|
|
if s.state.AccessToken != "" && time.Since(s.tokenReadAt) < 60*time.Second {
|
|
state := s.state
|
|
s.mu.Unlock()
|
|
return state, nil
|
|
}
|
|
s.mu.Unlock()
|
|
|
|
access, refresh, expiresAt, err := kiroReadSQLiteTokens(ctx, s.path)
|
|
if err != nil {
|
|
return s.state, nil // fall back to cached state
|
|
}
|
|
s.mu.Lock()
|
|
s.state.AccessToken = access
|
|
s.state.ExpiresAt = expiresAt
|
|
if refresh != "" {
|
|
s.state.RefreshToken = refresh
|
|
}
|
|
s.tokenReadAt = time.Now()
|
|
state := s.state
|
|
s.mu.Unlock()
|
|
return state, nil
|
|
}
|
|
|
|
// fullLoad performs a complete reload of all sqlite state including device
|
|
// registration (ClientID, ClientSecret). Used in Refresh() to pick up
|
|
// credentials rotated by another process (e.g. kiro-cli login).
|
|
func (s *sqliteKiroAuth) fullLoad(ctx context.Context) (kiroSQLiteState, error) {
|
|
state, err := kiroLoadSQLiteState(ctx, s.path)
|
|
if err != nil {
|
|
return kiroSQLiteState{}, err
|
|
}
|
|
s.initMu.Lock()
|
|
s.state = state
|
|
s.loaded = true
|
|
s.initMu.Unlock()
|
|
return state, nil
|
|
}
|
|
|
|
// tokenExpiringSoon checks if the token is about to expire, using instance
|
|
// jitter to spread refresh attempts across processes.
|
|
func (s *sqliteKiroAuth) tokenExpiringSoon(expiresAt string) bool {
|
|
if expiresAt == "" {
|
|
return false
|
|
}
|
|
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+s.jitter
|
|
}
|
|
|
|
// refreshFromSSOCache attempts to refresh using the AWS SSO cache as a
|
|
// fallback when the primary OIDC refresh token is invalid/expired.
|
|
func (s *sqliteKiroAuth) refreshFromSSOCache(ctx context.Context, previousAccess string) error {
|
|
cacheState, err := kiroLoadSSOCacheAuthState()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if strings.TrimSpace(cacheState.RefreshToken) == "" ||
|
|
strings.TrimSpace(cacheState.Region) == "" ||
|
|
strings.TrimSpace(cacheState.ClientID) == "" ||
|
|
strings.TrimSpace(cacheState.ClientSecret) == "" {
|
|
return fmt.Errorf("SSO cache does not contain enough data to refresh")
|
|
}
|
|
|
|
endpoint := fmt.Sprintf("https://oidc.%s.amazonaws.com/token", strings.TrimSpace(cacheState.Region))
|
|
client := newKiroOIDCClient(endpoint, nil)
|
|
|
|
resp, err := client.RefreshToken(ctx, kiroRefreshTokenRequest{
|
|
ClientID: cacheState.ClientID,
|
|
ClientSecret: cacheState.ClientSecret,
|
|
GrantType: "refresh_token",
|
|
RefreshToken: cacheState.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)
|
|
}
|
|
if err := kiroUpdateSQLiteTokens(ctx, s.path, newAccess, newRefresh, expiresAt); err != nil {
|
|
if s.sqliteHasNewUsableAccessToken(ctx, previousAccess) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
// Also update SSO cache file.
|
|
_ = kiroUpdateSSOCacheToken(cacheState.TokenPath, newAccess, newRefresh, expiresAt)
|
|
return nil
|
|
}
|
|
|
|
// sqliteHasNewUsableAccessToken retries reading the SQLite token to detect if
|
|
// another process refreshed during our attempt.
|
|
func (s *sqliteKiroAuth) sqliteHasNewUsableAccessToken(ctx context.Context, previousAccess string) bool {
|
|
const attempts = 3
|
|
const backoff = 500 * time.Millisecond
|
|
for i := range attempts {
|
|
if i > 0 {
|
|
select {
|
|
case <-ctx.Done():
|
|
return false
|
|
case <-time.After(backoff):
|
|
}
|
|
}
|
|
access, _, expiresAt, err := kiroReadSQLiteTokens(ctx, s.path)
|
|
if err != nil {
|
|
continue
|
|
}
|
|
if access != "" && access != strings.TrimSpace(previousAccess) && !s.tokenExpiringSoon(expiresAt) {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
const kiroTokenExpiryBuffer = 2 * time.Minute
|
|
|
|
// kiroTokenExpiryJitter is the maximum random offset added to kiroTokenExpiryBuffer.
|
|
const kiroTokenExpiryJitter = 2 * time.Minute
|
|
|
|
func kiroRandomJitter() time.Duration {
|
|
b := make([]byte, 4)
|
|
_, _ = rand.Read(b)
|
|
n := uint32(b[0]) | uint32(b[1])<<8 | uint32(b[2])<<16 | uint32(b[3])<<24
|
|
return time.Duration(n%uint32(kiroTokenExpiryJitter/time.Millisecond)) * time.Millisecond
|
|
}
|
|
|
|
// ── 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 kiroUpdateSQLiteSocialTokens(ctx context.Context, path, access, refresh, expiresAt, profileARN 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))
|
|
}
|
|
if p := strings.TrimSpace(profileARN); p != "" {
|
|
sb.WriteString(", '$.profile_arn', ")
|
|
sb.WriteString(kiroSQLiteQuote(p))
|
|
}
|
|
sb.WriteString(") WHERE key = 'kirocli:social: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
|
|
}
|
|
|
|
// sqliteBusyTimeout is the per-connection busy timeout passed to the sqlite3
|
|
// CLI via -cmd. This is session-scoped and does not modify the DB file.
|
|
const kiroSQLiteBusyTimeout = "5000"
|
|
|
|
// sqliteBusyRetries is the number of times we retry at the Go level when
|
|
// sqlite3 still returns a busy/locked error after its internal timeout.
|
|
const kiroSQLiteBusyRetries = 2
|
|
|
|
// sqliteBusyBackoff is the delay between Go-level retries.
|
|
const kiroSQLiteBusyBackoff = 200 * time.Millisecond
|
|
|
|
func kiroRunSQLite(ctx context.Context, args ...string) ([]byte, error) {
|
|
// Prepend busy timeout so sqlite3 retries internally on lock contention.
|
|
fullArgs := make([]string, 0, len(args)+2)
|
|
fullArgs = append(fullArgs, "-cmd", ".timeout "+kiroSQLiteBusyTimeout)
|
|
fullArgs = append(fullArgs, args...)
|
|
|
|
var out []byte
|
|
var err error
|
|
for attempt := range kiroSQLiteBusyRetries + 1 {
|
|
if attempt > 0 {
|
|
select {
|
|
case <-ctx.Done():
|
|
return out, ctx.Err()
|
|
case <-time.After(kiroSQLiteBusyBackoff):
|
|
}
|
|
}
|
|
cmd := exec.CommandContext(ctx, "sqlite3", fullArgs...)
|
|
out, err = cmd.CombinedOutput()
|
|
if err == nil || !kiroIsSQLiteBusy(out, err) {
|
|
return out, err
|
|
}
|
|
}
|
|
return out, err
|
|
}
|
|
|
|
// kiroIsSQLiteBusy checks whether the sqlite3 invocation failed due to lock
|
|
// contention. The sqlite3 CLI returns exit code 5 (SQLITE_BUSY) or 6
|
|
// (SQLITE_LOCKED).
|
|
func kiroIsSQLiteBusy(output []byte, err error) bool {
|
|
if exitErr, ok := err.(*exec.ExitError); ok {
|
|
code := exitErr.ExitCode()
|
|
if code == 5 || code == 6 {
|
|
return true
|
|
}
|
|
}
|
|
s := strings.ToLower(string(output))
|
|
return strings.Contains(s, "database is locked") ||
|
|
strings.Contains(s, "database table is locked")
|
|
}
|
|
|
|
func kiroSQLiteQuote(s string) string {
|
|
return "'" + strings.ReplaceAll(s, "'", "''") + "'"
|
|
}
|
|
|
|
// ── SSO cache fallback ────────────────────────────────────────────────────────
|
|
|
|
type kiroSSOCacheState struct {
|
|
TokenPath string
|
|
AccessToken string
|
|
RefreshToken string
|
|
ExpiresAt string
|
|
Region string
|
|
ClientID string
|
|
ClientSecret string
|
|
}
|
|
|
|
func kiroLoadSSOCacheAuthState() (kiroSSOCacheState, error) {
|
|
homeDir, err := os.UserHomeDir()
|
|
if err != nil {
|
|
return kiroSSOCacheState{}, fmt.Errorf("cannot determine home directory for SSO cache: %w", err)
|
|
}
|
|
cacheDir := filepath.Join(homeDir, ".aws", "sso", "cache")
|
|
tokenPath := filepath.Join(cacheDir, "kiro-auth-token-cli.json")
|
|
tokenBytes, err := os.ReadFile(tokenPath)
|
|
if err != nil {
|
|
return kiroSSOCacheState{}, fmt.Errorf("read SSO token cache %q: %w", tokenPath, err)
|
|
}
|
|
var tokenState struct {
|
|
AccessToken string `json:"accessToken"`
|
|
RefreshToken string `json:"refreshToken"`
|
|
ExpiresAt string `json:"expiresAt"`
|
|
Region string `json:"region"`
|
|
ClientIDHash string `json:"clientIdHash"`
|
|
}
|
|
if err := json.Unmarshal(tokenBytes, &tokenState); err != nil {
|
|
return kiroSSOCacheState{}, fmt.Errorf("decode SSO token cache %q: %w", tokenPath, err)
|
|
}
|
|
clientIDHash := strings.TrimSpace(tokenState.ClientIDHash)
|
|
if clientIDHash == "" {
|
|
return kiroSSOCacheState{}, fmt.Errorf("SSO token cache %q did not contain clientIdHash", tokenPath)
|
|
}
|
|
clientPath := filepath.Join(cacheDir, clientIDHash+".json")
|
|
clientBytes, err := os.ReadFile(clientPath)
|
|
if err != nil {
|
|
return kiroSSOCacheState{}, fmt.Errorf("read SSO client cache %q: %w", clientPath, err)
|
|
}
|
|
var clientState struct {
|
|
ClientID string `json:"clientId"`
|
|
ClientSecret string `json:"clientSecret"`
|
|
}
|
|
if err := json.Unmarshal(clientBytes, &clientState); err != nil {
|
|
return kiroSSOCacheState{}, fmt.Errorf("decode SSO client cache %q: %w", clientPath, err)
|
|
}
|
|
return kiroSSOCacheState{
|
|
TokenPath: tokenPath,
|
|
AccessToken: strings.TrimSpace(tokenState.AccessToken),
|
|
RefreshToken: strings.TrimSpace(tokenState.RefreshToken),
|
|
ExpiresAt: strings.TrimSpace(tokenState.ExpiresAt),
|
|
Region: strings.TrimSpace(tokenState.Region),
|
|
ClientID: strings.TrimSpace(clientState.ClientID),
|
|
ClientSecret: strings.TrimSpace(clientState.ClientSecret),
|
|
}, nil
|
|
}
|
|
|
|
func kiroUpdateSSOCacheToken(path, accessToken, refreshToken, expiresAt string) error {
|
|
accessToken = strings.TrimSpace(accessToken)
|
|
if accessToken == "" {
|
|
return fmt.Errorf("access token must not be empty")
|
|
}
|
|
body, err := os.ReadFile(path)
|
|
if err != nil {
|
|
return fmt.Errorf("read SSO token cache %q: %w", path, err)
|
|
}
|
|
var tokenState map[string]any
|
|
if err := json.Unmarshal(body, &tokenState); err != nil {
|
|
return fmt.Errorf("decode SSO token cache %q: %w", path, err)
|
|
}
|
|
tokenState["accessToken"] = accessToken
|
|
if strings.TrimSpace(refreshToken) != "" {
|
|
tokenState["refreshToken"] = strings.TrimSpace(refreshToken)
|
|
}
|
|
if strings.TrimSpace(expiresAt) != "" {
|
|
tokenState["expiresAt"] = strings.TrimSpace(expiresAt)
|
|
}
|
|
updated, err := json.MarshalIndent(tokenState, "", " ")
|
|
if err != nil {
|
|
return fmt.Errorf("encode SSO token cache %q: %w", path, err)
|
|
}
|
|
updated = append(updated, '\n')
|
|
info, err := os.Stat(path)
|
|
if err != nil {
|
|
return fmt.Errorf("stat SSO token cache %q: %w", path, err)
|
|
}
|
|
tmpPath := path + ".tmp"
|
|
if err := os.WriteFile(tmpPath, updated, info.Mode().Perm()); err != nil {
|
|
return fmt.Errorf("write SSO token cache temp %q: %w", tmpPath, err)
|
|
}
|
|
if err := os.Rename(tmpPath, path); err != nil {
|
|
os.Remove(tmpPath)
|
|
return fmt.Errorf("rename SSO token cache %q: %w", path, err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ── 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, jitter: kiroRandomJitter()}, nil
|
|
}
|
|
|
|
func kiroDefaultSQLiteKey() string {
|
|
return "sqlite://" + filepath.Join(util.UserDataHome(), "kiro-cli", "data.sqlite3")
|
|
}
|