ai_scheduler/internal/biz/llm_service/third_party/openai.go

200 lines
5.1 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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
// 对话场景采样参数:适度发散 + 高概率采样,降低「AI 腔」与复读感,更接近真人聊天
temperature := float32(0.8)
topP := float32(0.9)
req := openai.CreateResponseRequest{
Model: modelName,
Input: input,
Store: &store,
Temperature: &temperature,
TopP: &topP,
MaxOutputTokens: 800,
Stream: false,
Reasoning: &openai.ResponseReasoning{Effort: "none"},
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
}