239 lines
7.1 KiB
Go
239 lines
7.1 KiB
Go
package playground
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
// OpenAICompat drives the playground through any OpenAI-compatible chat
|
|
// completions API — Groq by default (https://api.groq.com/openai/v1), or xAI,
|
|
// or another provider — with plain net/http, so no SDK dependency is needed.
|
|
//
|
|
// Configured from PLAYGROUND_LLM_BASE_URL / _API_KEY / _MODEL (see config).
|
|
// The model comes from that setting, not from the agent's registry pin: the
|
|
// registry holds Claude ids for AI_engine, which this provider cannot serve.
|
|
type OpenAICompat struct {
|
|
BaseURL string
|
|
APIKey string
|
|
Model string
|
|
HTTP *http.Client
|
|
}
|
|
|
|
// NewOpenAICompat returns a client with a bounded HTTP timeout.
|
|
func NewOpenAICompat(baseURL, apiKey, model string) *OpenAICompat {
|
|
return &OpenAICompat{
|
|
BaseURL: strings.TrimRight(baseURL, "/"),
|
|
APIKey: apiKey,
|
|
Model: model,
|
|
HTTP: &http.Client{Timeout: 90 * time.Second},
|
|
}
|
|
}
|
|
|
|
// ModelName is what the trace reports as the model that answered.
|
|
func (o *OpenAICompat) ModelName() string { return o.Model }
|
|
|
|
// ── Wire types (OpenAI chat completions) ────────────────────────────────────
|
|
|
|
type oaFunctionCall struct {
|
|
Name string `json:"name"`
|
|
Arguments string `json:"arguments"`
|
|
}
|
|
|
|
type oaToolCall struct {
|
|
ID string `json:"id"`
|
|
Type string `json:"type"`
|
|
Function oaFunctionCall `json:"function"`
|
|
}
|
|
|
|
type oaMessage struct {
|
|
Role string `json:"role"`
|
|
Content *string `json:"content"`
|
|
ToolCalls []oaToolCall `json:"tool_calls,omitempty"`
|
|
ToolCallID string `json:"tool_call_id,omitempty"`
|
|
}
|
|
|
|
type oaTool struct {
|
|
Type string `json:"type"`
|
|
Function struct {
|
|
Name string `json:"name"`
|
|
Description string `json:"description"`
|
|
Parameters json.RawMessage `json:"parameters"`
|
|
} `json:"function"`
|
|
}
|
|
|
|
type oaRequest struct {
|
|
Model string `json:"model"`
|
|
Messages []oaMessage `json:"messages"`
|
|
Tools []oaTool `json:"tools,omitempty"`
|
|
MaxCompletionTokens int64 `json:"max_completion_tokens,omitempty"`
|
|
}
|
|
|
|
type oaResponse struct {
|
|
Choices []struct {
|
|
Message oaMessage `json:"message"`
|
|
FinishReason string `json:"finish_reason"`
|
|
} `json:"choices"`
|
|
Usage struct {
|
|
PromptTokens int64 `json:"prompt_tokens"`
|
|
CompletionTokens int64 `json:"completion_tokens"`
|
|
} `json:"usage"`
|
|
Error *struct {
|
|
Message string `json:"message"`
|
|
} `json:"error"`
|
|
}
|
|
|
|
func strp(s string) *string { return &s }
|
|
|
|
// objectSchema makes sure a tool's parameters are an object schema with a
|
|
// properties map — OpenAI-style APIs reject `{"type":"object"}` alone on some
|
|
// models, and several registry tools have open schemas.
|
|
func objectSchema(raw json.RawMessage) json.RawMessage {
|
|
var m map[string]any
|
|
if len(raw) == 0 || json.Unmarshal(raw, &m) != nil || m == nil {
|
|
m = map[string]any{}
|
|
}
|
|
m["type"] = "object"
|
|
if _, ok := m["properties"]; !ok {
|
|
m["properties"] = map[string]any{}
|
|
}
|
|
b, _ := json.Marshal(m)
|
|
return b
|
|
}
|
|
|
|
// toWire converts the playground conversation into chat-completions messages.
|
|
func toWire(req Request) oaRequest {
|
|
out := oaRequest{MaxCompletionTokens: req.MaxTokens}
|
|
if req.System != "" {
|
|
out.Messages = append(out.Messages, oaMessage{Role: "system", Content: strp(req.System)})
|
|
}
|
|
for _, t := range req.Turns {
|
|
switch {
|
|
case t.Role == "assistant":
|
|
msg := oaMessage{Role: "assistant"}
|
|
var texts []string
|
|
for _, b := range t.Assistant {
|
|
switch {
|
|
case b.Type == "text" && b.Text != "":
|
|
texts = append(texts, b.Text)
|
|
case b.Type == "tool_use" && b.ToolUse != nil:
|
|
args := string(b.ToolUse.Input)
|
|
if args == "" {
|
|
args = "{}"
|
|
}
|
|
msg.ToolCalls = append(msg.ToolCalls, oaToolCall{
|
|
ID: b.ToolUse.ID, Type: "function",
|
|
Function: oaFunctionCall{Name: b.ToolUse.Name, Arguments: args},
|
|
})
|
|
}
|
|
}
|
|
if len(texts) > 0 {
|
|
msg.Content = strp(strings.Join(texts, "\n\n"))
|
|
}
|
|
out.Messages = append(out.Messages, msg)
|
|
case len(t.Results) > 0:
|
|
for _, r := range t.Results {
|
|
out.Messages = append(out.Messages, oaMessage{Role: "tool", ToolCallID: r.ToolUseID, Content: strp(r.Content)})
|
|
}
|
|
default:
|
|
out.Messages = append(out.Messages, oaMessage{Role: "user", Content: strp(t.Text)})
|
|
}
|
|
}
|
|
for _, td := range req.Tools {
|
|
var tool oaTool
|
|
tool.Type = "function"
|
|
tool.Function.Name = td.Name
|
|
tool.Function.Description = td.Description
|
|
tool.Function.Parameters = objectSchema(td.InputSchema)
|
|
out.Tools = append(out.Tools, tool)
|
|
}
|
|
return out
|
|
}
|
|
|
|
// fromWire converts one chat-completions response into a playground Reply.
|
|
func fromWire(resp oaResponse) (Reply, error) {
|
|
if len(resp.Choices) == 0 {
|
|
return Reply{}, errors.New("model returned no choices")
|
|
}
|
|
ch := resp.Choices[0]
|
|
r := Reply{InputTokens: resp.Usage.PromptTokens, OutputTokens: resp.Usage.CompletionTokens}
|
|
if ch.Message.Content != nil && strings.TrimSpace(*ch.Message.Content) != "" {
|
|
r.Blocks = append(r.Blocks, Block{Type: "text", Text: *ch.Message.Content})
|
|
}
|
|
for _, tc := range ch.Message.ToolCalls {
|
|
input := json.RawMessage(tc.Function.Arguments)
|
|
if !json.Valid(input) {
|
|
input = json.RawMessage(`{}`)
|
|
}
|
|
r.Blocks = append(r.Blocks, Block{Type: "tool_use", ToolUse: &ToolUse{ID: tc.ID, Name: tc.Function.Name, Input: input}})
|
|
}
|
|
switch ch.FinishReason {
|
|
case "tool_calls":
|
|
r.StopReason = "tool_use"
|
|
case "length":
|
|
r.StopReason = "max_tokens"
|
|
default:
|
|
r.StopReason = "end_turn"
|
|
}
|
|
// Some providers report "stop" while still returning tool calls; the calls win.
|
|
if len(ch.Message.ToolCalls) > 0 {
|
|
r.StopReason = "tool_use"
|
|
}
|
|
return r, nil
|
|
}
|
|
|
|
// Next performs one chat-completions call.
|
|
func (o *OpenAICompat) Next(ctx context.Context, req Request) (Reply, error) {
|
|
body := toWire(req)
|
|
body.Model = o.Model
|
|
payload, err := json.Marshal(body)
|
|
if err != nil {
|
|
return Reply{}, fmt.Errorf("encode request: %w", err)
|
|
}
|
|
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, o.BaseURL+"/chat/completions", bytes.NewReader(payload))
|
|
if err != nil {
|
|
return Reply{}, fmt.Errorf("build request: %w", err)
|
|
}
|
|
httpReq.Header.Set("Authorization", "Bearer "+o.APIKey)
|
|
httpReq.Header.Set("Content-Type", "application/json")
|
|
|
|
res, err := o.HTTP.Do(httpReq)
|
|
if err != nil {
|
|
return Reply{}, fmt.Errorf("model request: %w", err)
|
|
}
|
|
defer res.Body.Close()
|
|
raw, _ := io.ReadAll(io.LimitReader(res.Body, 4<<20))
|
|
|
|
var out oaResponse
|
|
_ = json.Unmarshal(raw, &out)
|
|
if res.StatusCode/100 != 2 {
|
|
msg := strings.TrimSpace(string(raw))
|
|
if out.Error != nil && out.Error.Message != "" {
|
|
msg = out.Error.Message
|
|
}
|
|
if len(msg) > 300 {
|
|
msg = msg[:300]
|
|
}
|
|
return Reply{}, &ProviderError{Status: res.StatusCode, Message: msg}
|
|
}
|
|
return fromWire(out)
|
|
}
|
|
|
|
// ProviderError is a non-2xx answer from the model provider. The status lets
|
|
// the controller tell "rate limited" (429, common on free tiers) from a
|
|
// genuine failure. The message never contains the API key.
|
|
type ProviderError struct {
|
|
Status int
|
|
Message string
|
|
}
|
|
|
|
func (e *ProviderError) Error() string {
|
|
return fmt.Sprintf("model provider answered %d: %s", e.Status, e.Message)
|
|
}
|