106 lines
3.4 KiB
Go
106 lines
3.4 KiB
Go
package backend
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"testing"
|
|
)
|
|
|
|
func TestKiroAPIError_IsThrottled(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
err kiroAPIError
|
|
expect bool
|
|
}{
|
|
{"429 status", kiroAPIError{StatusCode: http.StatusTooManyRequests}, true},
|
|
{"ThrottlingException code", kiroAPIError{Code: "ThrottlingException", StatusCode: 400}, true},
|
|
{"401 not throttled", kiroAPIError{StatusCode: http.StatusUnauthorized}, false},
|
|
{"other code", kiroAPIError{Code: "ValidationException", StatusCode: 400}, false},
|
|
}
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
if got := tt.err.isThrottled(); got != tt.expect {
|
|
t.Errorf("isThrottled() = %v, want %v", got, tt.expect)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestKiroStripToolBlocks(t *testing.T) {
|
|
tests := []struct {
|
|
name string
|
|
messages []Message
|
|
want int // expected number of messages after stripping
|
|
desc string
|
|
}{
|
|
{
|
|
name: "no tool content passes through unchanged",
|
|
messages: []Message{
|
|
{Role: "system", Content: "You are helpful."},
|
|
{Role: "user", Content: "Hello"},
|
|
{Role: "assistant", Content: "Hi there!"},
|
|
},
|
|
want: 3,
|
|
desc: "all messages preserved",
|
|
},
|
|
{
|
|
name: "tool messages are removed",
|
|
messages: []Message{
|
|
{Role: "user", Content: "Read a file"},
|
|
{Role: "assistant", Content: "", ToolCalls: []ToolCall{{ID: "tc1", Name: "file_read", Arguments: json.RawMessage(`{"path":"/tmp/x"}`)}}},
|
|
{Role: "tool", Content: "file contents here", ToolCallID: "tc1"},
|
|
{Role: "assistant", Content: "The file says..."},
|
|
},
|
|
want: 2, // user + final assistant (tool-only assistant and tool msg dropped)
|
|
desc: "tool-only assistant and tool result dropped",
|
|
},
|
|
{
|
|
name: "assistant with text and tool calls keeps text",
|
|
messages: []Message{
|
|
{Role: "user", Content: "Do something"},
|
|
{Role: "assistant", Content: "Let me check that.", ToolCalls: []ToolCall{{ID: "tc1", Name: "shell", Arguments: json.RawMessage(`{}`)}}},
|
|
{Role: "tool", Content: "ok", ToolCallID: "tc1"},
|
|
{Role: "assistant", Content: "Done!"},
|
|
},
|
|
want: 3, // user + assistant(text only) + final assistant
|
|
desc: "assistant with text preserved without tool calls",
|
|
},
|
|
{
|
|
name: "multiple consecutive tool messages all stripped",
|
|
messages: []Message{
|
|
{Role: "user", Content: "Run stuff"},
|
|
{Role: "assistant", ToolCalls: []ToolCall{
|
|
{ID: "tc1", Name: "a", Arguments: json.RawMessage(`{}`)},
|
|
{ID: "tc2", Name: "b", Arguments: json.RawMessage(`{}`)},
|
|
}},
|
|
{Role: "tool", Content: "r1", ToolCallID: "tc1"},
|
|
{Role: "tool", Content: "r2", ToolCallID: "tc2"},
|
|
{Role: "user", Content: "Thanks"},
|
|
},
|
|
want: 2, // user + user("Thanks")
|
|
desc: "tool-only assistant + tool results all dropped",
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
got := kiroStripToolBlocks(tt.messages)
|
|
if len(got) != tt.want {
|
|
t.Errorf("%s: got %d messages, want %d", tt.desc, len(got), tt.want)
|
|
for i, m := range got {
|
|
t.Logf(" [%d] role=%s content=%q toolCalls=%d", i, m.Role, m.Content, len(m.ToolCalls))
|
|
}
|
|
}
|
|
// Verify no remaining messages have tool content.
|
|
for i, m := range got {
|
|
if m.Role == "tool" {
|
|
t.Errorf("message[%d] still has role=tool", i)
|
|
}
|
|
if len(m.ToolCalls) > 0 {
|
|
t.Errorf("message[%d] still has %d tool calls", i, len(m.ToolCalls))
|
|
}
|
|
}
|
|
})
|
|
}
|
|
}
|