nomadcode/services/core/internal/adapters/openai/client.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
}