498 lines
13 KiB
Go
498 lines
13 KiB
Go
package openai
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/url"
|
|
"path"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/nomadcode/nomadcode-core/internal/model"
|
|
)
|
|
|
|
const (
|
|
defaultTimeoutSec = 900
|
|
responsesPath = "/v1/responses"
|
|
)
|
|
|
|
type Config struct {
|
|
BaseURL string
|
|
APIKey string
|
|
Model string
|
|
ContextSize int
|
|
TimeoutSec int
|
|
ModelResponsesStream bool
|
|
}
|
|
|
|
type Client struct {
|
|
cfg Config
|
|
httpClient *http.Client
|
|
logger *slog.Logger
|
|
}
|
|
|
|
func NewClient(cfg Config, logger *slog.Logger) *Client {
|
|
if cfg.TimeoutSec <= 0 {
|
|
cfg.TimeoutSec = defaultTimeoutSec
|
|
}
|
|
return &Client{
|
|
cfg: cfg,
|
|
httpClient: &http.Client{Timeout: time.Duration(cfg.TimeoutSec) * time.Second},
|
|
logger: logger,
|
|
}
|
|
}
|
|
|
|
func (c *Client) Generate(ctx context.Context, input model.GenerateInput) (model.GenerateResult, error) {
|
|
endpoint, err := responsesURL(c.cfg.BaseURL)
|
|
if err != nil {
|
|
return model.GenerateResult{}, err
|
|
}
|
|
|
|
modelName := strings.TrimSpace(input.Model)
|
|
if modelName == "" {
|
|
modelName = c.cfg.Model
|
|
}
|
|
if strings.TrimSpace(modelName) == "" {
|
|
return model.GenerateResult{}, fmt.Errorf("model name is required")
|
|
}
|
|
if strings.TrimSpace(input.Input) == "" {
|
|
return model.GenerateResult{}, fmt.Errorf("model input is required")
|
|
}
|
|
|
|
streamRequested := c.cfg.ModelResponsesStream
|
|
respResult, err := c.executeGenerate(ctx, input, endpoint, modelName, streamRequested)
|
|
return respResult, err
|
|
}
|
|
|
|
func (c *Client) executeGenerate(ctx context.Context, input model.GenerateInput, endpoint, modelName string, stream bool) (model.GenerateResult, error) {
|
|
reqBody := responsesRequest{
|
|
Model: modelName,
|
|
Input: input.Input,
|
|
Instructions: input.Instructions,
|
|
Metadata: buildRequestMetadata(input),
|
|
Stream: stream,
|
|
Temperature: input.Temperature,
|
|
TopP: input.TopP,
|
|
MaxOutputTokens: input.MaxOutputTokens,
|
|
}
|
|
if c.cfg.ContextSize > 0 {
|
|
reqBody.Options = &responsesOptions{NumCtx: c.cfg.ContextSize}
|
|
}
|
|
body, err := json.Marshal(reqBody)
|
|
if err != nil {
|
|
return model.GenerateResult{}, err
|
|
}
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body))
|
|
if err != nil {
|
|
return model.GenerateResult{}, err
|
|
}
|
|
req.Header.Set("Content-Type", "application/json")
|
|
if c.cfg.APIKey != "" {
|
|
req.Header.Set("Authorization", "Bearer "+c.cfg.APIKey)
|
|
}
|
|
if stream {
|
|
req.Header.Set("Accept", "text/event-stream")
|
|
}
|
|
|
|
if c.logger != nil {
|
|
c.logger.Info(
|
|
"model responses request",
|
|
"endpoint", endpoint,
|
|
"model", modelName,
|
|
"num_ctx", c.cfg.ContextSize,
|
|
"stream", stream,
|
|
)
|
|
}
|
|
|
|
resp, err := c.httpClient.Do(req)
|
|
if err != nil {
|
|
if stream && isStreamUnsupportedError(err) {
|
|
if input.OnProgress != nil {
|
|
input.OnProgress(model.GenerateProgress{
|
|
Mode: "stream_unsupported",
|
|
Reason: err.Error(),
|
|
LastEventTime: time.Now().UTC(),
|
|
})
|
|
}
|
|
return c.executeGenerate(ctx, input, endpoint, modelName, false)
|
|
}
|
|
return model.GenerateResult{}, err
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
|
errRes := responseError(resp.StatusCode, raw)
|
|
if stream && isStreamUnsupportedError(errRes) {
|
|
if input.OnProgress != nil {
|
|
input.OnProgress(model.GenerateProgress{
|
|
Mode: "stream_unsupported",
|
|
Reason: errRes.Error(),
|
|
LastEventTime: time.Now().UTC(),
|
|
})
|
|
}
|
|
return c.executeGenerate(ctx, input, endpoint, modelName, false)
|
|
}
|
|
return model.GenerateResult{}, errRes
|
|
}
|
|
|
|
if stream {
|
|
reader := bufio.NewReader(resp.Body)
|
|
var finalResponse responsesResponse
|
|
var textBuilder strings.Builder
|
|
var currentEvent string
|
|
for {
|
|
line, err := reader.ReadString('\n')
|
|
if err != nil {
|
|
if err == io.EOF {
|
|
break
|
|
}
|
|
return model.GenerateResult{}, err
|
|
}
|
|
line = strings.TrimSpace(line)
|
|
if line == "" {
|
|
continue
|
|
}
|
|
if strings.HasPrefix(line, "event: ") {
|
|
currentEvent = strings.TrimPrefix(line, "event: ")
|
|
continue
|
|
}
|
|
if !strings.HasPrefix(line, "data: ") {
|
|
continue
|
|
}
|
|
dataStr := strings.TrimPrefix(line, "data: ")
|
|
if dataStr == "[DONE]" {
|
|
break
|
|
}
|
|
|
|
var ev responsesStreamEvent
|
|
errEv := json.Unmarshal([]byte(dataStr), &ev)
|
|
if errEv == nil && (ev.Type != "" || ev.Error != nil || ev.Response != nil) {
|
|
eventType := ev.Type
|
|
if eventType == "" {
|
|
eventType = currentEvent
|
|
}
|
|
|
|
switch eventType {
|
|
case "response.output_text.delta":
|
|
textBuilder.WriteString(ev.Delta)
|
|
if input.OnProgress != nil {
|
|
input.OnProgress(model.GenerateProgress{
|
|
Mode: "streaming",
|
|
Reason: "receiving stream chunks",
|
|
LastEventTime: time.Now().UTC(),
|
|
})
|
|
}
|
|
case "response.completed":
|
|
if ev.Response != nil {
|
|
finalResponse.ID = ev.Response.ID
|
|
finalResponse.Model = ev.Response.Model
|
|
if ev.Response.Usage != nil {
|
|
finalResponse.Usage = ev.Response.Usage
|
|
}
|
|
}
|
|
case "error":
|
|
if ev.Error != nil {
|
|
return model.GenerateResult{}, fmt.Errorf("responses API stream error: %s", ev.Error.Message)
|
|
}
|
|
return model.GenerateResult{}, fmt.Errorf("responses API stream error: unknown error")
|
|
}
|
|
} else {
|
|
var chunk responsesResponse
|
|
errChunk := json.Unmarshal([]byte(dataStr), &chunk)
|
|
if errChunk == nil {
|
|
if chunk.ID != "" {
|
|
finalResponse.ID = chunk.ID
|
|
}
|
|
if chunk.Model != "" {
|
|
finalResponse.Model = chunk.Model
|
|
}
|
|
if chunk.Usage != nil {
|
|
if finalResponse.Usage == nil {
|
|
finalResponse.Usage = &responsesUsage{}
|
|
}
|
|
if chunk.Usage.InputTokens != 0 {
|
|
finalResponse.Usage.InputTokens = chunk.Usage.InputTokens
|
|
}
|
|
if chunk.Usage.OutputTokens != 0 {
|
|
finalResponse.Usage.OutputTokens = chunk.Usage.OutputTokens
|
|
}
|
|
if chunk.Usage.TotalTokens != 0 {
|
|
finalResponse.Usage.TotalTokens = chunk.Usage.TotalTokens
|
|
}
|
|
if chunk.Usage.PromptTokens != 0 {
|
|
finalResponse.Usage.PromptTokens = chunk.Usage.PromptTokens
|
|
}
|
|
if chunk.Usage.CompletionTokens != 0 {
|
|
finalResponse.Usage.CompletionTokens = chunk.Usage.CompletionTokens
|
|
}
|
|
}
|
|
chunkText := chunk.text()
|
|
if chunkText != "" {
|
|
textBuilder.WriteString(chunkText)
|
|
}
|
|
if input.OnProgress != nil {
|
|
input.OnProgress(model.GenerateProgress{
|
|
Mode: "streaming",
|
|
Reason: "receiving stream chunks",
|
|
LastEventTime: time.Now().UTC(),
|
|
})
|
|
}
|
|
} else {
|
|
return model.GenerateResult{}, fmt.Errorf("failed to parse SSE data: %s (errEv: %v, errChunk: %v)", dataStr, errEv, errChunk)
|
|
}
|
|
}
|
|
currentEvent = ""
|
|
}
|
|
|
|
finalResponse.OutputText = textBuilder.String()
|
|
raw, _ := json.Marshal(finalResponse)
|
|
|
|
usage := model.Usage{}
|
|
if finalResponse.Usage != nil {
|
|
usage.InputTokens = firstNonZero(finalResponse.Usage.InputTokens, finalResponse.Usage.PromptTokens)
|
|
usage.OutputTokens = firstNonZero(finalResponse.Usage.OutputTokens, finalResponse.Usage.CompletionTokens)
|
|
usage.TotalTokens = finalResponse.Usage.TotalTokens
|
|
if usage.TotalTokens == 0 {
|
|
usage.TotalTokens = usage.InputTokens + usage.OutputTokens
|
|
}
|
|
}
|
|
|
|
return model.GenerateResult{
|
|
ID: finalResponse.ID,
|
|
Model: firstNonEmpty(finalResponse.Model, modelName),
|
|
Text: finalResponse.text(),
|
|
Usage: usage,
|
|
Raw: json.RawMessage(raw),
|
|
}, nil
|
|
}
|
|
|
|
if input.OnProgress != nil {
|
|
input.OnProgress(model.GenerateProgress{
|
|
Mode: "non_streaming",
|
|
Reason: "non-streaming mode",
|
|
LastEventTime: time.Now().UTC(),
|
|
})
|
|
}
|
|
|
|
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
|
if err != nil {
|
|
return model.GenerateResult{}, err
|
|
}
|
|
|
|
var parsed responsesResponse
|
|
if err := json.Unmarshal(raw, &parsed); err != nil {
|
|
return model.GenerateResult{}, err
|
|
}
|
|
|
|
usage := model.Usage{}
|
|
if parsed.Usage != nil {
|
|
usage.InputTokens = firstNonZero(parsed.Usage.InputTokens, parsed.Usage.PromptTokens)
|
|
usage.OutputTokens = firstNonZero(parsed.Usage.OutputTokens, parsed.Usage.CompletionTokens)
|
|
usage.TotalTokens = parsed.Usage.TotalTokens
|
|
if usage.TotalTokens == 0 {
|
|
usage.TotalTokens = usage.InputTokens + usage.OutputTokens
|
|
}
|
|
}
|
|
|
|
return model.GenerateResult{
|
|
ID: parsed.ID,
|
|
Model: firstNonEmpty(parsed.Model, modelName),
|
|
Text: parsed.text(),
|
|
Usage: usage,
|
|
Raw: json.RawMessage(raw),
|
|
}, nil
|
|
}
|
|
|
|
type responsesRequest struct {
|
|
Model string `json:"model"`
|
|
Input string `json:"input"`
|
|
Instructions string `json:"instructions,omitempty"`
|
|
Metadata map[string]any `json:"metadata,omitempty"`
|
|
Stream bool `json:"stream"`
|
|
Temperature *float64 `json:"temperature,omitempty"`
|
|
TopP *float64 `json:"top_p,omitempty"`
|
|
MaxOutputTokens int `json:"max_output_tokens,omitempty"`
|
|
Options *responsesOptions `json:"options,omitempty"`
|
|
}
|
|
|
|
type responsesOptions struct {
|
|
NumCtx int `json:"num_ctx,omitempty"`
|
|
}
|
|
|
|
type responsesResponse struct {
|
|
ID string `json:"id"`
|
|
Model string `json:"model"`
|
|
OutputText string `json:"output_text"`
|
|
Output []responsesOutput `json:"output"`
|
|
Usage *responsesUsage `json:"usage"`
|
|
}
|
|
|
|
type responsesOutput struct {
|
|
Type string `json:"type"`
|
|
Text string `json:"text"`
|
|
Content []responsesContent `json:"content"`
|
|
}
|
|
|
|
type responsesContent struct {
|
|
Type string `json:"type"`
|
|
Text string `json:"text"`
|
|
}
|
|
|
|
type responsesUsage struct {
|
|
InputTokens int `json:"input_tokens"`
|
|
OutputTokens int `json:"output_tokens"`
|
|
TotalTokens int `json:"total_tokens"`
|
|
PromptTokens int `json:"prompt_tokens"`
|
|
CompletionTokens int `json:"completion_tokens"`
|
|
}
|
|
|
|
func (r responsesResponse) text() string {
|
|
if r.OutputText != "" {
|
|
return r.OutputText
|
|
}
|
|
|
|
var b strings.Builder
|
|
for _, output := range r.Output {
|
|
if output.Text != "" {
|
|
b.WriteString(output.Text)
|
|
}
|
|
for _, content := range output.Content {
|
|
if content.Text != "" {
|
|
b.WriteString(content.Text)
|
|
}
|
|
}
|
|
}
|
|
return b.String()
|
|
}
|
|
|
|
type errorResponse struct {
|
|
Error struct {
|
|
Message string `json:"message"`
|
|
Type string `json:"type"`
|
|
Code string `json:"code"`
|
|
} `json:"error"`
|
|
}
|
|
|
|
func responseError(status int, raw []byte) error {
|
|
var parsed errorResponse
|
|
if err := json.Unmarshal(raw, &parsed); err == nil && parsed.Error.Message != "" {
|
|
return fmt.Errorf("responses API request failed: status %d: %s", status, parsed.Error.Message)
|
|
}
|
|
msg := strings.TrimSpace(string(raw))
|
|
if msg == "" {
|
|
msg = http.StatusText(status)
|
|
}
|
|
return fmt.Errorf("responses API request failed: status %d: %s", status, msg)
|
|
}
|
|
|
|
func responsesURL(base string) (string, error) {
|
|
base = strings.TrimSpace(base)
|
|
if base == "" {
|
|
return "", fmt.Errorf("model base URL is required")
|
|
}
|
|
|
|
parsed, err := url.Parse(base)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if parsed.Scheme == "" || parsed.Host == "" {
|
|
return "", fmt.Errorf("model base URL must include scheme and host")
|
|
}
|
|
|
|
cleanPath := strings.TrimRight(parsed.Path, "/")
|
|
switch {
|
|
case strings.HasSuffix(cleanPath, "/v1/responses"):
|
|
parsed.Path = cleanPath
|
|
case strings.HasSuffix(cleanPath, "/v1"):
|
|
parsed.Path = path.Join(cleanPath, "responses")
|
|
default:
|
|
parsed.Path = path.Join(cleanPath, responsesPath)
|
|
}
|
|
return parsed.String(), nil
|
|
}
|
|
|
|
// buildRequestMetadata merges flat string metadata with typed WorkspaceMetadata.
|
|
// WorkspaceMetadata wins on key collision ("workspace" key).
|
|
func buildRequestMetadata(input model.GenerateInput) map[string]any {
|
|
if len(input.Metadata) == 0 && input.WorkspaceMetadata == nil {
|
|
return nil
|
|
}
|
|
merged := make(map[string]any, len(input.Metadata)+1)
|
|
for k, v := range input.Metadata {
|
|
// IOP Responses contract rejects ambiguous call-origin metadata.
|
|
if k == "source" {
|
|
continue
|
|
}
|
|
merged[k] = v
|
|
}
|
|
if input.WorkspaceMetadata != nil {
|
|
if path := strings.TrimSpace(input.WorkspaceMetadata.Path); path != "" {
|
|
merged["workspace"] = path
|
|
}
|
|
}
|
|
if len(merged) == 0 {
|
|
return nil
|
|
}
|
|
return merged
|
|
}
|
|
|
|
func firstNonZero(values ...int) int {
|
|
for _, value := range values {
|
|
if value != 0 {
|
|
return value
|
|
}
|
|
}
|
|
return 0
|
|
}
|
|
|
|
func firstNonEmpty(values ...string) string {
|
|
for _, value := range values {
|
|
if value != "" {
|
|
return value
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
type responsesStreamEvent struct {
|
|
Type string `json:"type"`
|
|
Delta string `json:"delta,omitempty"`
|
|
Response *struct {
|
|
ID string `json:"id,omitempty"`
|
|
Model string `json:"model,omitempty"`
|
|
Usage *responsesUsage `json:"usage,omitempty"`
|
|
} `json:"response,omitempty"`
|
|
Error *struct {
|
|
Message string `json:"message,omitempty"`
|
|
Type string `json:"type,omitempty"`
|
|
Code string `json:"code,omitempty"`
|
|
} `json:"error,omitempty"`
|
|
}
|
|
|
|
func isStreamUnsupportedError(err error) bool {
|
|
if err == nil {
|
|
return false
|
|
}
|
|
errStr := strings.ToLower(err.Error())
|
|
signals := []string{
|
|
"stream",
|
|
"streaming",
|
|
"unsupported",
|
|
"not supported",
|
|
"unsupported_parameter",
|
|
}
|
|
for _, sig := range signals {
|
|
if strings.Contains(errStr, sig) {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|