初始化
This commit is contained in:
@@ -0,0 +1,141 @@
|
||||
package apimart
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"juhe-factory/api/internal/provider"
|
||||
)
|
||||
|
||||
type ChatResult struct {
|
||||
Content string
|
||||
Raw json.RawMessage
|
||||
FinishReason string
|
||||
InputTokens int64
|
||||
OutputTokens int64
|
||||
TotalTokens int64
|
||||
UsageRaw json.RawMessage
|
||||
}
|
||||
|
||||
func (c Client) Chat(ctx context.Context, baseURL, apiKey string, payload map[string]any) (string, json.RawMessage, error) {
|
||||
result, err := c.ChatWithUsage(ctx, baseURL, apiKey, payload)
|
||||
return result.Content, result.Raw, err
|
||||
}
|
||||
|
||||
func (c Client) ChatWithUsage(ctx context.Context, baseURL, apiKey string, payload map[string]any) (ChatResult, error) {
|
||||
var result ChatResult
|
||||
requestPayload := make(map[string]any, len(payload)+1)
|
||||
for key, value := range payload {
|
||||
requestPayload[key] = value
|
||||
}
|
||||
requestPayload["stream"] = false
|
||||
body, err := json.Marshal(requestPayload)
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimRight(baseURL, "/")+"/chat/completions", bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
request.Header.Set("Authorization", "Bearer "+apiKey)
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
response, err := c.HTTPClient.Do(request)
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
defer response.Body.Close()
|
||||
data, err := io.ReadAll(io.LimitReader(response.Body, 16*1024*1024))
|
||||
if err != nil {
|
||||
return result, err
|
||||
}
|
||||
if response.StatusCode < 200 || response.StatusCode >= 300 {
|
||||
result.Raw = append(json.RawMessage(nil), data...)
|
||||
return result, &provider.HTTPError{StatusCode: response.StatusCode, Body: safeBody(data), RetryAfter: response.Header.Get("Retry-After")}
|
||||
}
|
||||
type chatChoice struct {
|
||||
Message struct {
|
||||
Content any `json:"content"`
|
||||
} `json:"message"`
|
||||
FinishReason string `json:"finish_reason"`
|
||||
}
|
||||
var decoded struct {
|
||||
Choices []chatChoice `json:"choices"`
|
||||
Usage json.RawMessage `json:"usage"`
|
||||
Data struct {
|
||||
Choices []chatChoice `json:"choices"`
|
||||
Usage json.RawMessage `json:"usage"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
result.Raw = append(json.RawMessage(nil), data...)
|
||||
return result, fmt.Errorf("文本模型响应不是有效 JSON: %w;上游响应: %s", err, safeBody(data))
|
||||
}
|
||||
result.Raw = append(json.RawMessage(nil), data...)
|
||||
result.UsageRaw = decoded.Usage
|
||||
choices := decoded.Choices
|
||||
if len(choices) == 0 {
|
||||
choices = decoded.Data.Choices
|
||||
if len(result.UsageRaw) == 0 {
|
||||
result.UsageRaw = decoded.Data.Usage
|
||||
}
|
||||
}
|
||||
if len(choices) == 0 {
|
||||
return result, fmt.Errorf("文本模型响应缺少 choices;上游响应: %s", safeBody(data))
|
||||
}
|
||||
result.FinishReason = choices[0].FinishReason
|
||||
if len(result.UsageRaw) > 0 {
|
||||
var usage map[string]any
|
||||
if json.Unmarshal(result.UsageRaw, &usage) == nil {
|
||||
result.InputTokens = intValue(usage["prompt_tokens"], usage["input_tokens"])
|
||||
result.OutputTokens = intValue(usage["completion_tokens"], usage["output_tokens"])
|
||||
result.TotalTokens = intValue(usage["total_tokens"])
|
||||
if result.TotalTokens == 0 {
|
||||
result.TotalTokens = result.InputTokens + result.OutputTokens
|
||||
}
|
||||
}
|
||||
}
|
||||
switch content := choices[0].Message.Content.(type) {
|
||||
case string:
|
||||
result.Content = content
|
||||
return result, nil
|
||||
case []any:
|
||||
parts := make([]string, 0)
|
||||
for _, item := range content {
|
||||
if block, ok := item.(map[string]any); ok {
|
||||
if text, ok := block["text"].(string); ok {
|
||||
parts = append(parts, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
result.Content = strings.Join(parts, "\n")
|
||||
return result, nil
|
||||
default:
|
||||
return result, errors.New("文本模型响应内容无效")
|
||||
}
|
||||
}
|
||||
|
||||
func intValue(values ...any) int64 {
|
||||
for _, value := range values {
|
||||
switch number := value.(type) {
|
||||
case float64:
|
||||
if number > 0 {
|
||||
return int64(number)
|
||||
}
|
||||
case int64:
|
||||
if number > 0 {
|
||||
return number
|
||||
}
|
||||
case json.Number:
|
||||
if parsed, err := number.Int64(); err == nil && parsed > 0 {
|
||||
return parsed
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
Reference in New Issue
Block a user