192 lines
4.7 KiB
Go
192 lines
4.7 KiB
Go
package third_party
|
||
|
||
import (
|
||
"context"
|
||
"net/http"
|
||
"sync"
|
||
"time"
|
||
|
||
"github.com/gofiber/fiber/v2/log"
|
||
openai "github.com/sashabaranov/go-openai"
|
||
)
|
||
|
||
type OpenAi struct {
|
||
mapClient map[string]*openai.Client
|
||
mu sync.RWMutex
|
||
}
|
||
|
||
var (
|
||
openAi *OpenAi
|
||
openAiOnce sync.Once
|
||
)
|
||
|
||
// NewOpenAi 返回单例
|
||
func NewOpenAi() *OpenAi {
|
||
openAiOnce.Do(func() {
|
||
openAi = &OpenAi{
|
||
mapClient: make(map[string]*openai.Client),
|
||
}
|
||
})
|
||
return openAi
|
||
}
|
||
|
||
// getClient 按 key + baseURL 缓存 client
|
||
func (o *OpenAi) getClient(key string, baseURL string) *openai.Client {
|
||
cacheKey := key + "|" + baseURL
|
||
|
||
o.mu.RLock()
|
||
if c, ok := o.mapClient[cacheKey]; ok {
|
||
o.mu.RUnlock()
|
||
return c
|
||
}
|
||
o.mu.RUnlock()
|
||
|
||
cfg := openai.DefaultConfig(key)
|
||
if baseURL != "" {
|
||
cfg.BaseURL = baseURL
|
||
}
|
||
cfg.HTTPClient = &http.Client{Timeout: 2 * time.Minute}
|
||
|
||
client := openai.NewClientWithConfig(cfg)
|
||
|
||
o.mu.Lock()
|
||
o.mapClient[cacheKey] = client
|
||
o.mu.Unlock()
|
||
|
||
return client
|
||
}
|
||
|
||
// CreateResponse 对应 Hsyq.CreateResponse
|
||
// id 不为空时走 PreviousResponseID 续写对话
|
||
func (o *OpenAi) CreateResponse(
|
||
ctx context.Context,
|
||
key string,
|
||
modelName string,
|
||
input string,
|
||
id string,
|
||
) (*openai.CreateResponseResponse, error) {
|
||
req := openai.CreateResponseRequest{
|
||
Model: modelName,
|
||
Input: input,
|
||
}
|
||
|
||
if len(id) != 0 {
|
||
req.PreviousResponseID = id
|
||
}
|
||
|
||
resp, err := o.getClient(key, "").CreateResponse(ctx, req)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if resp.Usage != nil {
|
||
log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.InputTokens, "输出:", resp.Usage.OutputTokens)
|
||
}
|
||
return &resp, nil
|
||
}
|
||
|
||
// CreateResponseMessages 对应 Hsyq.CreateResponse 的多消息版本
|
||
// json_object格式输出;响应需存储后才能通过 PreviousResponseID 续写
|
||
// id 不为空时走 PreviousResponseID 续写对话
|
||
func (o *OpenAi) CreateResponseMessages(
|
||
ctx context.Context,
|
||
key string,
|
||
url string,
|
||
modelName string,
|
||
input []openai.ResponseInputMessage,
|
||
id string,
|
||
) (*openai.CreateResponseResponse, error) {
|
||
store := true
|
||
req := openai.CreateResponseRequest{
|
||
Model: modelName,
|
||
Input: input,
|
||
Store: &store,
|
||
Text: &openai.ResponseTextConfig{
|
||
Format: &openai.ResponseTextFormat{
|
||
Type: "json_object",
|
||
},
|
||
},
|
||
}
|
||
|
||
if len(id) != 0 {
|
||
req.PreviousResponseID = id
|
||
}
|
||
|
||
resp, err := o.getClient(key, url).CreateResponse(ctx, req)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if resp.Usage != nil {
|
||
log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.InputTokens, "输出:", resp.Usage.OutputTokens)
|
||
}
|
||
return &resp, nil
|
||
}
|
||
|
||
// RequestResponsesJson 对应 Hsyq.RequestHsyqJson
|
||
// 通过 Text.Format 指定 JSON 输出,直接返回文字内容
|
||
func (o *OpenAi) RequestResponsesJson(
|
||
ctx context.Context,
|
||
key string,
|
||
url string,
|
||
modelName string,
|
||
input string,
|
||
) (string, error) {
|
||
req := openai.CreateResponseRequest{
|
||
Model: modelName,
|
||
Input: input,
|
||
Text: &openai.ResponseTextConfig{
|
||
Format: &openai.ResponseTextFormat{
|
||
Type: "json_object",
|
||
},
|
||
},
|
||
}
|
||
|
||
resp, err := o.getClient(key, url).CreateResponse(ctx, req)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
if resp.Usage != nil {
|
||
log.Info("token用量:", resp.Usage.TotalTokens)
|
||
}
|
||
|
||
return resp.GetOutputText(), nil
|
||
}
|
||
|
||
// Chat 基础对话(保留 Chat Completions 兼容)
|
||
func (o *OpenAi) Chat(ctx context.Context, key string, modelName string, prompt []openai.ChatCompletionMessage) (openai.ChatCompletionResponse, error) {
|
||
req := openai.ChatCompletionRequest{
|
||
Model: modelName,
|
||
Messages: prompt,
|
||
Stream: false,
|
||
}
|
||
resp, err := o.getClient(key, "").CreateChatCompletion(ctx, req)
|
||
if err != nil {
|
||
return openai.ChatCompletionResponse{}, err
|
||
}
|
||
log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens, "输出:", resp.Usage.CompletionTokens)
|
||
return resp, nil
|
||
}
|
||
|
||
// ChatWithRequest 自定义 Chat 请求
|
||
func (o *OpenAi) ChatWithRequest(ctx context.Context, key string, request openai.ChatCompletionRequest) (openai.ChatCompletionResponse, error) {
|
||
resp, err := o.getClient(key, "").CreateChatCompletion(ctx, request)
|
||
if err != nil {
|
||
return openai.ChatCompletionResponse{}, err
|
||
}
|
||
log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens, "输出:", resp.Usage.CompletionTokens)
|
||
return resp, nil
|
||
}
|
||
|
||
// CreateEmbedding 向量化
|
||
func (o *OpenAi) CreateEmbedding(ctx context.Context, key string, modelName string, input []string) (openai.EmbeddingResponse, error) {
|
||
req := openai.EmbeddingRequest{
|
||
Model: openai.EmbeddingModel(modelName),
|
||
Input: input,
|
||
}
|
||
resp, err := o.getClient(key, "").CreateEmbeddings(ctx, req)
|
||
if err != nil {
|
||
return openai.EmbeddingResponse{}, err
|
||
}
|
||
log.Info("token用量:", resp.Usage.TotalTokens, "输入:", resp.Usage.PromptTokens)
|
||
return resp, nil
|
||
}
|