improved tool debugging
This commit is contained in:
@@ -1,122 +1,217 @@
|
||||
package sidekick
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// LLMClient represents an abstraction over the LLM provider
|
||||
type LLMClient interface {
|
||||
Call(ctx context.Context, model ModelConfig, messages []Message, tools map[string]ToolDefinition) (Message, error)
|
||||
Call(ctx context.Context, model ModelConfig, messages []Message,
|
||||
tools map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error)
|
||||
}
|
||||
|
||||
// standardLLMClient is an OpenAI-compatible implementation
|
||||
type standardLLMClient struct{}
|
||||
|
||||
func (s *standardLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, ts map[string]ToolDefinition) (Message, error) {
|
||||
client := &http.Client{
|
||||
Timeout: model.Timeout,
|
||||
}
|
||||
func (s *standardLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message,
|
||||
ts map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) {
|
||||
client := &http.Client{
|
||||
Timeout: model.Timeout,
|
||||
}
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"model": model.ModelID,
|
||||
"messages": msgs,
|
||||
"temperature": model.Temperature,
|
||||
}
|
||||
if model.MaxTokens > 0 {
|
||||
reqBody["max_tokens"] = model.MaxTokens
|
||||
}
|
||||
reqBody := map[string]interface{}{
|
||||
"model": model.ModelID,
|
||||
"messages": msgs,
|
||||
"temperature": model.Temperature,
|
||||
}
|
||||
if model.MaxTokens > 0 {
|
||||
reqBody["max_tokens"] = model.MaxTokens
|
||||
}
|
||||
|
||||
if len(ts) > 0 {
|
||||
var tools []map[string]interface{}
|
||||
for _, t := range ts {
|
||||
tools = append(tools, t.ConvertToOpenAITool())
|
||||
}
|
||||
reqBody["tools"] = tools
|
||||
}
|
||||
if len(ts) > 0 {
|
||||
var tools []map[string]interface{}
|
||||
for _, t := range ts {
|
||||
tools = append(tools, t.ConvertToOpenAITool())
|
||||
}
|
||||
reqBody["tools"] = tools
|
||||
}
|
||||
|
||||
jsonBody, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
if streamCallback != nil {
|
||||
reqBody["stream"] = true
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", model.Endpoint+"/chat/completions", bytes.NewBuffer(jsonBody))
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
jsonBody, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+model.APIKey)
|
||||
req, err := http.NewRequestWithContext(ctx, "POST",
|
||||
model.Endpoint+"/chat/completions", bytes.NewBuffer(jsonBody))
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("HTTP request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+model.APIKey)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to read response body: %w", err)
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("HTTP request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return Message{}, fmt.Errorf("API returned error (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return Message{}, fmt.Errorf("API returned error (%d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var openAIResp struct {
|
||||
Choices []struct {
|
||||
Message Message `json:"message"`
|
||||
} `json:"choices"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error,omitempty"`
|
||||
}
|
||||
if streamCallback != nil {
|
||||
return s.handleStream(resp.Body, streamCallback, reasoningCallback)
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(body, &openAIResp); err != nil {
|
||||
return Message{}, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return Message{}, fmt.Errorf("failed to read response body: %w", err)
|
||||
}
|
||||
|
||||
if openAIResp.Error != nil {
|
||||
return Message{}, fmt.Errorf("API error: %s", openAIResp.Error.Message)
|
||||
}
|
||||
var openAIResp struct {
|
||||
Choices []struct {
|
||||
Message Message `json:"message"`
|
||||
} `json:"choices"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
if len(openAIResp.Choices) == 0 {
|
||||
return Message{}, fmt.Errorf("API returned no choices")
|
||||
}
|
||||
if err := json.Unmarshal(body, &openAIResp); err != nil {
|
||||
return Message{}, fmt.Errorf("failed to unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
return openAIResp.Choices[0].Message, nil
|
||||
if openAIResp.Error != nil {
|
||||
return Message{}, fmt.Errorf("API error: %s", openAIResp.Error.Message)
|
||||
}
|
||||
|
||||
if len(openAIResp.Choices) == 0 {
|
||||
return Message{}, fmt.Errorf("API returned no choices")
|
||||
}
|
||||
|
||||
return openAIResp.Choices[0].Message, nil
|
||||
}
|
||||
|
||||
var defaultLLMClient LLMClient = &standardLLMClient{}
|
||||
|
||||
// SetLLMClient allows overriding the LLM client (e.g. for testing)
|
||||
func SetLLMClient(client LLMClient) {
|
||||
defaultLLMClient = client
|
||||
defaultLLMClient = client
|
||||
}
|
||||
|
||||
// CallLLM invokes the configured LLM client
|
||||
func CallLLM(ctx context.Context, model ModelConfig, msgs []Message, tools map[string]ToolDefinition) (Message, error) {
|
||||
return defaultLLMClient.Call(ctx, model, msgs, tools)
|
||||
func CallLLM(ctx context.Context, model ModelConfig, msgs []Message,
|
||||
tools map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) {
|
||||
return defaultLLMClient.Call(ctx, model, msgs, tools, streamCallback, reasoningCallback)
|
||||
}
|
||||
|
||||
// mockLLMClient is used as a fallback if no actual HTTP client is injected.
|
||||
type mockLLMClient struct {
|
||||
responses []Message
|
||||
calls int
|
||||
responses []Message
|
||||
calls int
|
||||
}
|
||||
|
||||
func (m *mockLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message, ts map[string]ToolDefinition) (Message, error) {
|
||||
func (m *mockLLMClient) Call(ctx context.Context, model ModelConfig, msgs []Message,
|
||||
ts map[string]ToolDefinition, streamCallback, reasoningCallback func(string)) (Message, error) {
|
||||
if m.calls < len(m.responses) {
|
||||
resp := m.responses[m.calls]
|
||||
m.calls++
|
||||
if streamCallback != nil {
|
||||
streamCallback(resp.Content)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
return Message{
|
||||
Role: MessageRoleAssistant,
|
||||
msg := Message{
|
||||
Role: MessageRoleAssistant,
|
||||
Content: `{"action": "RESPOND", "params": {"response": "Mock LLM Response"}}`,
|
||||
}, nil
|
||||
}
|
||||
if streamCallback != nil {
|
||||
streamCallback(msg.Content)
|
||||
}
|
||||
return msg, nil
|
||||
}
|
||||
|
||||
func (s *standardLLMClient) handleStream(body io.ReadCloser, streamCallback,
|
||||
reasoningCallback func(string)) (Message, error) {
|
||||
defer body.Close()
|
||||
var fullContent strings.Builder
|
||||
var finalMsg Message
|
||||
finalMsg.Role = MessageRoleAssistant
|
||||
|
||||
type deltaToolCall struct {
|
||||
Index int `json:"index"`
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Function struct {
|
||||
Name string `json:"name"`
|
||||
Arguments string `json:"arguments"`
|
||||
} `json:"function"`
|
||||
}
|
||||
|
||||
toolCallsMap := make(map[int]*ToolCall)
|
||||
maxIndex := -1
|
||||
|
||||
scanner := bufio.NewScanner(body)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
if !strings.HasPrefix(line, "data: ") {
|
||||
continue
|
||||
}
|
||||
data := strings.TrimPrefix(line, "data: ")
|
||||
if data == "[DONE]" {
|
||||
break
|
||||
}
|
||||
|
||||
var chunk struct {
|
||||
Choices []struct {
|
||||
Delta struct {
|
||||
Content string `json:"content"`
|
||||
ToolCalls []deltaToolCall `json:"tool_calls"`
|
||||
} `json:"delta"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(data), &chunk); err == nil && len(chunk.Choices) > 0 {
|
||||
delta := chunk.Choices[0].Delta
|
||||
if delta.Content != "" {
|
||||
fullContent.WriteString(delta.Content)
|
||||
streamCallback(delta.Content)
|
||||
}
|
||||
for _, tc := range delta.ToolCalls {
|
||||
if _, exists := toolCallsMap[tc.Index]; !exists {
|
||||
toolCallsMap[tc.Index] = &ToolCall{ID: tc.ID, Type: tc.Type}
|
||||
if tc.Index > maxIndex {
|
||||
maxIndex = tc.Index
|
||||
}
|
||||
}
|
||||
toolCallsMap[tc.Index].Function.Name += tc.Function.Name
|
||||
toolCallsMap[tc.Index].Function.Arguments += tc.Function.Arguments
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return Message{}, fmt.Errorf("error reading stream: %w", err)
|
||||
}
|
||||
|
||||
finalMsg.Content = fullContent.String()
|
||||
if maxIndex >= 0 {
|
||||
for i := 0; i <= maxIndex; i++ {
|
||||
if tc, ok := toolCallsMap[i]; ok {
|
||||
finalMsg.ToolCalls = append(finalMsg.ToolCalls, *tc)
|
||||
}
|
||||
}
|
||||
}
|
||||
return finalMsg, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user