576 lines
18 KiB
Go
576 lines
18 KiB
Go
package service
|
||
|
||
import (
|
||
"bufio"
|
||
"bytes"
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
"io"
|
||
"math"
|
||
"net/http"
|
||
"sort"
|
||
"strconv"
|
||
"strings"
|
||
"time"
|
||
|
||
"rag-local/common"
|
||
"rag-local/kb/consts"
|
||
"rag-local/kb/dao"
|
||
"rag-local/kb/model/domain"
|
||
"rag-local/kb/model/entity"
|
||
|
||
eembedding "github.com/cloudwego/eino/components/embedding"
|
||
emodel "github.com/cloudwego/eino/components/model"
|
||
eretriever "github.com/cloudwego/eino/components/retriever"
|
||
"github.com/cloudwego/eino/schema"
|
||
"github.com/gogf/gf/v2/errors/gerror"
|
||
"github.com/gogf/gf/v2/frame/g"
|
||
)
|
||
|
||
var httpClient = &http.Client{Timeout: 2 * time.Minute}
|
||
|
||
// ---------- OpenAI 兼容 HTTP 组件 ----------
|
||
|
||
type openAIMessage struct {
|
||
Role string `json:"role"`
|
||
Content string `json:"content"`
|
||
}
|
||
|
||
type openAIChatResponse struct {
|
||
Choices []struct {
|
||
Message openAIMessage `json:"message"`
|
||
FinishReason string `json:"finish_reason"`
|
||
} `json:"choices"`
|
||
Usage *schema.TokenUsage `json:"usage"`
|
||
}
|
||
|
||
type openAIStreamChunk struct {
|
||
Choices []struct {
|
||
Delta openAIMessage `json:"delta"`
|
||
FinishReason string `json:"finish_reason"`
|
||
} `json:"choices"`
|
||
}
|
||
|
||
// OpenAIChatModel 基于 OpenAI 兼容 /chat/completions 接口的对话模型,实现 eino model.ChatModel
|
||
type OpenAIChatModel struct {
|
||
cfg *entity.ModelConfig
|
||
}
|
||
|
||
func NewOpenAIChatModel(cfg *entity.ModelConfig) *OpenAIChatModel {
|
||
return &OpenAIChatModel{cfg: cfg}
|
||
}
|
||
|
||
func (m *OpenAIChatModel) Generate(ctx context.Context, input []*schema.Message, opts ...emodel.Option) (*schema.Message, error) {
|
||
payload := map[string]any{
|
||
"model": m.cfg.ModelName,
|
||
"messages": buildOpenAIMessages(input),
|
||
"stream": false,
|
||
}
|
||
body, err := postOpenAI(ctx, m.cfg, m.endpoint("/chat/completions"), payload)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
var resp openAIChatResponse
|
||
if err := json.Unmarshal(body, &resp); err != nil {
|
||
return nil, gerror.Wrap(err, "解析模型响应失败")
|
||
}
|
||
if len(resp.Choices) == 0 {
|
||
return nil, gerror.New("模型返回空响应")
|
||
}
|
||
choice := resp.Choices[0]
|
||
msg := &schema.Message{Role: schema.Assistant, Content: choice.Message.Content}
|
||
if choice.FinishReason != "" {
|
||
msg.ResponseMeta = &schema.ResponseMeta{FinishReason: choice.FinishReason, Usage: resp.Usage}
|
||
}
|
||
return msg, nil
|
||
}
|
||
|
||
func (m *OpenAIChatModel) Stream(ctx context.Context, input []*schema.Message, opts ...emodel.Option) (*schema.StreamReader[*schema.Message], error) {
|
||
payload := map[string]any{
|
||
"model": m.cfg.ModelName,
|
||
"messages": buildOpenAIMessages(input),
|
||
"stream": true,
|
||
}
|
||
reader, writer := schema.Pipe[*schema.Message](16)
|
||
go func() {
|
||
defer writer.Close()
|
||
body, err := postOpenAIStream(ctx, m.cfg, m.endpoint("/chat/completions"), payload)
|
||
if err != nil {
|
||
writer.Send(nil, err)
|
||
return
|
||
}
|
||
defer body.Close()
|
||
br := bufio.NewReader(body)
|
||
for {
|
||
line, err := br.ReadBytes('\n')
|
||
text := strings.TrimSpace(string(line))
|
||
if strings.HasPrefix(text, "data:") {
|
||
data := strings.TrimSpace(strings.TrimPrefix(text, "data:"))
|
||
if data == "[DONE]" {
|
||
break
|
||
}
|
||
var chunk openAIStreamChunk
|
||
if json.Unmarshal([]byte(data), &chunk) == nil && len(chunk.Choices) > 0 {
|
||
if delta := chunk.Choices[0].Delta.Content; delta != "" {
|
||
if closed := writer.Send(&schema.Message{Role: schema.Assistant, Content: delta}, nil); closed {
|
||
return
|
||
}
|
||
}
|
||
}
|
||
}
|
||
if err != nil {
|
||
break
|
||
}
|
||
}
|
||
}()
|
||
return reader, nil
|
||
}
|
||
|
||
func (m *OpenAIChatModel) BindTools(tools []*schema.ToolInfo) error {
|
||
return nil
|
||
}
|
||
|
||
func (m *OpenAIChatModel) endpoint(path string) string {
|
||
return strings.TrimRight(m.cfg.EndpointUrl, "/") + path
|
||
}
|
||
|
||
// OpenAIEmbedder 基于 OpenAI 兼容 /embeddings 接口的向量模型,实现 eino embedding.Embedder
|
||
type OpenAIEmbedder struct {
|
||
cfg *entity.ModelConfig
|
||
}
|
||
|
||
func NewOpenAIEmbedder(cfg *entity.ModelConfig) *OpenAIEmbedder {
|
||
return &OpenAIEmbedder{cfg: cfg}
|
||
}
|
||
|
||
// Dim 配置的向量维度(vec0 表建表维度需一致)
|
||
func (e *OpenAIEmbedder) Dim() int {
|
||
if e.cfg.Dimension > 0 {
|
||
return e.cfg.Dimension
|
||
}
|
||
return consts.DefaultEmbeddingDim
|
||
}
|
||
|
||
// EmbedStrings 内部分批请求(dashscope 单次最多 20 条),返回与 texts 同序的向量
|
||
func (e *OpenAIEmbedder) EmbedStrings(ctx context.Context, texts []string, opts ...eembedding.Option) ([][]float64, error) {
|
||
out := make([][]float64, len(texts))
|
||
for start := 0; start < len(texts); start += consts.EmbedBatchSize {
|
||
end := min(start+consts.EmbedBatchSize, len(texts))
|
||
payload := map[string]any{"model": e.cfg.ModelName, "input": texts[start:end]}
|
||
body, err := postOpenAI(ctx, e.cfg, strings.TrimRight(e.cfg.EndpointUrl, "/")+"/embeddings", payload)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
var resp struct {
|
||
Data []struct {
|
||
Embedding []float64 `json:"embedding"`
|
||
Index int `json:"index"`
|
||
} `json:"data"`
|
||
}
|
||
if err := json.Unmarshal(body, &resp); err != nil {
|
||
return nil, gerror.Wrap(err, "解析 embedding 响应失败")
|
||
}
|
||
if len(resp.Data) == 0 {
|
||
return nil, gerror.New("embedding 接口返回空数据")
|
||
}
|
||
for _, d := range resp.Data {
|
||
if d.Index >= 0 && d.Index < end-start {
|
||
out[start+d.Index] = d.Embedding
|
||
}
|
||
}
|
||
}
|
||
return out, nil
|
||
}
|
||
|
||
// BuildChatModel 按配置 id 构建对话模型
|
||
func BuildChatModel(ctx context.Context, cfgId int64) (*OpenAIChatModel, error) {
|
||
cfg, err := dao.ModelConfig.GetOne(ctx, cfgId)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if cfg == nil {
|
||
return nil, gerror.New("模型配置不存在")
|
||
}
|
||
if cfg.ModelType != consts.ModelTypeChat {
|
||
return nil, gerror.New("该模型配置不是对话模型(model_type=chat)")
|
||
}
|
||
if cfg.EndpointUrl == "" || cfg.ModelName == "" {
|
||
return nil, gerror.New("模型配置缺少 endpoint_url 或 model_name")
|
||
}
|
||
return NewOpenAIChatModel(cfg), nil
|
||
}
|
||
|
||
// BuildEmbedder 按配置 id 构建向量模型
|
||
func BuildEmbedder(ctx context.Context, cfgId int64) (*OpenAIEmbedder, error) {
|
||
cfg, err := dao.ModelConfig.GetOne(ctx, cfgId)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
if cfg == nil {
|
||
return nil, gerror.New("模型配置不存在")
|
||
}
|
||
if cfg.ModelType != consts.ModelTypeEmbedding {
|
||
return nil, gerror.New("该模型配置不是向量模型(model_type=embedding)")
|
||
}
|
||
if cfg.EndpointUrl == "" || cfg.ModelName == "" {
|
||
return nil, gerror.New("模型配置缺少 endpoint_url 或 model_name")
|
||
}
|
||
return NewOpenAIEmbedder(cfg), nil
|
||
}
|
||
|
||
// HybridRetriever 混合检索器:向量 KNN + FTS5 BM25,RRF 融合后可选 LLM 重排,实现 eino retriever.Retriever
|
||
type HybridRetriever struct {
|
||
embedder eembedding.Embedder
|
||
datasetId int64
|
||
reranker *OpenAIChatModel // LLM 重排器(默认对话模型),为 nil 时跳过重排走 RRF 顺序
|
||
}
|
||
|
||
func NewHybridRetriever(embedder eembedding.Embedder, datasetId int64, reranker *OpenAIChatModel) *HybridRetriever {
|
||
return &HybridRetriever{embedder: embedder, datasetId: datasetId, reranker: reranker}
|
||
}
|
||
|
||
func (r *HybridRetriever) Retrieve(ctx context.Context, query string, opts ...eretriever.Option) ([]*schema.Document, error) {
|
||
o := eretriever.GetCommonOptions(nil, opts...)
|
||
topK := consts.HybridTopK
|
||
if o.TopK != nil && *o.TopK > 0 {
|
||
topK = *o.TopK
|
||
}
|
||
|
||
scores := make(map[int64]float64)
|
||
srcs := make(map[int64][]string)
|
||
|
||
if r.embedder != nil {
|
||
vecs, err := r.embedder.EmbedStrings(ctx, []string{query})
|
||
if err != nil {
|
||
g.Log().Warningf(ctx, "query embed failed: %v", err)
|
||
} else if len(vecs) > 0 {
|
||
hits, err := dao.Chunk.VecSearch(ctx, r.datasetId, domain.VecJsonF64(vecs[0]), consts.VectorTopK)
|
||
if err != nil {
|
||
g.Log().Warningf(ctx, "vec search failed: %v", err)
|
||
} else {
|
||
for i, h := range hits {
|
||
scores[h.ChunkId] += 1 / (float64(consts.RrfK) + float64(i) + 1)
|
||
srcs[h.ChunkId] = append(srcs[h.ChunkId], "vector")
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
ftsHits, err := dao.Chunk.FtsSearch(ctx, r.datasetId, common.TokenizeQuery(query), consts.FtsTopK)
|
||
if err != nil {
|
||
g.Log().Warningf(ctx, "fts search failed: %v", err)
|
||
} else {
|
||
for i, h := range ftsHits {
|
||
scores[h.ChunkId] += 1 / (float64(consts.RrfK) + float64(i) + 1)
|
||
srcs[h.ChunkId] = append(srcs[h.ChunkId], "fts")
|
||
}
|
||
}
|
||
|
||
items := make([]scoredChunk, 0, len(scores))
|
||
for id, s := range scores {
|
||
items = append(items, scoredChunk{id: id, score: s, sources: srcs[id]})
|
||
}
|
||
sort.Slice(items, func(i, j int) bool { return items[i].score > items[j].score })
|
||
if len(items) > consts.RerankTopK {
|
||
items = items[:consts.RerankTopK]
|
||
}
|
||
if r.reranker != nil && len(items) > 0 {
|
||
before := len(items)
|
||
if scores, err := r.rerankByLLM(ctx, query, items); err == nil {
|
||
for i := range items {
|
||
items[i].score = scores[items[i].id]
|
||
}
|
||
sort.Slice(items, func(i, j int) bool { return items[i].score > items[j].score })
|
||
// 门槛作用在重排语义分上:最高分条目必留;其余需同时满足相对比例与绝对下限,
|
||
// 防止重排器对泛化条款给出宽松低分(如 3 分)也能进引用
|
||
maxScore := items[0].score
|
||
if maxScore > 0 {
|
||
floor := math.Max(maxScore*consts.RerankKeepRatio, consts.RerankMinScore)
|
||
keep := items[:0]
|
||
for i, it := range items {
|
||
if i == 0 || it.score >= floor {
|
||
keep = append(keep, it)
|
||
}
|
||
}
|
||
items = keep
|
||
g.Log().Infof(ctx, "rerank done: %d candidates → %d kept (max %.1f, floor %.1f)",
|
||
before, len(items), maxScore, floor)
|
||
}
|
||
} else {
|
||
g.Log().Warningf(ctx, "rerank failed, fallback to rrf order: %v", err)
|
||
}
|
||
}
|
||
if len(items) > topK {
|
||
items = items[:topK]
|
||
}
|
||
|
||
docs := make([]*schema.Document, 0, len(items))
|
||
for _, it := range items {
|
||
chunk, err := dao.Chunk.GetOne(ctx, it.id)
|
||
if err != nil || chunk == nil {
|
||
continue
|
||
}
|
||
docs = append(docs, &schema.Document{
|
||
ID: strconv.FormatInt(chunk.Id, 10),
|
||
Content: chunk.Content,
|
||
MetaData: map[string]any{
|
||
"chunk_id": chunk.Id,
|
||
"document_id": chunk.DocumentId,
|
||
"seq": chunk.Seq,
|
||
"score": it.score,
|
||
"sources": it.sources,
|
||
},
|
||
})
|
||
}
|
||
return docs, nil
|
||
}
|
||
|
||
// scoredChunk RRF 融合后的候选条目(重排后 score 字段替换为语义分)
|
||
type scoredChunk struct {
|
||
id int64
|
||
score float64
|
||
sources []string
|
||
}
|
||
|
||
// rerankByLLM 用默认对话模型对候选分块打分(0-10,JSON 输出),返回 chunk_id → 相关性分;
|
||
// 任何失败(调用/解析/空结果)返回错误,由调用方回退 RRF 排序
|
||
func (r *HybridRetriever) rerankByLLM(ctx context.Context, query string, items []scoredChunk) (map[int64]float64, error) {
|
||
var sb strings.Builder
|
||
sb.WriteString("你是检索重排器。请评估每个候选段落与用户问题的相关性,为每个候选输出 0-10 的相关性分数(10=高度相关,0=完全不相关)。\n\n用户问题:\n")
|
||
sb.WriteString(query)
|
||
sb.WriteString("\n\n候选段落:\n")
|
||
for i, it := range items {
|
||
chunk, err := dao.Chunk.GetOne(ctx, it.id)
|
||
if err != nil || chunk == nil {
|
||
continue
|
||
}
|
||
content := chunk.Content
|
||
if rs := []rune(content); len(rs) > consts.RerankMaxChars {
|
||
content = string(rs[:consts.RerankMaxChars])
|
||
}
|
||
sb.WriteString(fmt.Sprintf("[%d] %s\n", i+1, content))
|
||
}
|
||
sb.WriteString("\n只输出 JSON,不要其他内容:{\"scores\":{\"1\":8,\"2\":3}}")
|
||
msg, err := r.reranker.Generate(ctx, []*schema.Message{{Role: schema.User, Content: sb.String()}})
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
// 模型可能包裹 ```json 代码块,截取首尾花括号之间的 JSON 对象
|
||
content := msg.Content
|
||
if i, j := strings.Index(content, "{"), strings.LastIndex(content, "}"); i >= 0 && j > i {
|
||
content = content[i : j+1]
|
||
}
|
||
var resp struct {
|
||
Scores map[string]float64 `json:"scores"`
|
||
}
|
||
if err := json.Unmarshal([]byte(content), &resp); err != nil {
|
||
return nil, gerror.Wrap(err, "解析重排结果失败: "+msg.Content)
|
||
}
|
||
scores := make(map[int64]float64, len(resp.Scores))
|
||
for k, v := range resp.Scores {
|
||
idx, err := strconv.ParseInt(k, 10, 64)
|
||
if err != nil || idx < 1 || idx > int64(len(items)) {
|
||
continue
|
||
}
|
||
scores[items[idx-1].id] = v
|
||
}
|
||
if len(scores) == 0 {
|
||
return nil, gerror.New("重排结果为空")
|
||
}
|
||
return scores, nil
|
||
}
|
||
|
||
// ---------- RAG 问答工作流 ----------
|
||
|
||
var ChatService = &chatService{}
|
||
|
||
type chatService struct{}
|
||
|
||
// MaxHistoryRounds 携带进模型的历史对话轮数(每条消息算一条,含用户与助手)
|
||
const MaxHistoryRounds = 10
|
||
|
||
// Ask RAG 问答工作流:混合检索 → 组装提示(含引用编号)→ 对话模型流式生成。
|
||
// history 需已包含最新一条用户问题;onCitations 在检索完成后先于流式输出回调;onDelta 接收增量文本,均可为 nil。
|
||
func (s *chatService) Ask(ctx context.Context, datasetId int64, question string, history []*schema.Message, onCitations func([]domain.Citation), onDelta func(string)) (string, []domain.Citation, error) {
|
||
docs, err := s.retrieve(ctx, datasetId, question)
|
||
if err != nil {
|
||
return "", nil, err
|
||
}
|
||
citations := buildCitations(docs)
|
||
if onCitations != nil {
|
||
onCitations(citations)
|
||
}
|
||
|
||
triples, err := KgRelationService.GraphEnhance(ctx, datasetId, question)
|
||
if err != nil {
|
||
g.Log().Warningf(ctx, "graph enhance failed: %v", err)
|
||
}
|
||
|
||
defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat)
|
||
if err != nil {
|
||
return "", nil, err
|
||
}
|
||
if defaultChatModel <= 0 {
|
||
return "", nil, gerror.New("请先在设置中为对话模型设置默认")
|
||
}
|
||
model, err := BuildChatModel(ctx, defaultChatModel)
|
||
if err != nil {
|
||
return "", nil, err
|
||
}
|
||
|
||
msgs := make([]*schema.Message, 0, len(history)+1)
|
||
msgs = append(msgs, &schema.Message{Role: schema.System, Content: buildSystemPrompt(citations, triples)})
|
||
if start := len(history) - MaxHistoryRounds*2; start > 0 {
|
||
history = history[start:]
|
||
}
|
||
msgs = append(msgs, history...)
|
||
|
||
sr, err := model.Stream(ctx, msgs)
|
||
if err != nil {
|
||
return "", nil, gerror.Wrap(err, "调用对话模型失败")
|
||
}
|
||
defer sr.Close()
|
||
var full strings.Builder
|
||
for {
|
||
m, err := sr.Recv()
|
||
if errors.Is(err, io.EOF) {
|
||
break
|
||
}
|
||
if err != nil {
|
||
return "", nil, gerror.Wrap(err, "流式输出中断")
|
||
}
|
||
full.WriteString(m.Content)
|
||
if onDelta != nil {
|
||
onDelta(m.Content)
|
||
}
|
||
}
|
||
return full.String(), citations, nil
|
||
}
|
||
|
||
// retrieve 构建数据集绑定的混合检索器并执行检索(无 embedding 配置时仅全文)
|
||
func (s *chatService) retrieve(ctx context.Context, datasetId int64, question string) ([]*schema.Document, error) {
|
||
var emb eembedding.Embedder
|
||
if cfgId, err := dao.Dataset.GetEmbeddingCfgId(ctx, datasetId); err == nil && cfgId > 0 {
|
||
if em, err := BuildEmbedder(ctx, cfgId); err == nil {
|
||
emb = em
|
||
} else {
|
||
g.Log().Warningf(ctx, "build embedder failed, retrieve fts only: %v", err)
|
||
}
|
||
}
|
||
// LLM 重排器:默认对话模型;构建失败仅跳过重排,不影响检索主流程
|
||
var reranker *OpenAIChatModel
|
||
if defaultChatModel, err := dao.ModelConfig.GetDefault(ctx, consts.ModelTypeChat); err == nil && defaultChatModel > 0 {
|
||
if m, err := BuildChatModel(ctx, defaultChatModel); err == nil {
|
||
reranker = m
|
||
} else {
|
||
g.Log().Warningf(ctx, "build reranker failed, skip rerank: %v", err)
|
||
}
|
||
}
|
||
return NewHybridRetriever(emb, datasetId, reranker).Retrieve(ctx, question)
|
||
}
|
||
|
||
// buildCitations 从检索结果生成引用列表(编号从 1 开始,与提示词 [编号] 对应)
|
||
func buildCitations(docs []*schema.Document) []domain.Citation {
|
||
cits := make([]domain.Citation, 0, len(docs))
|
||
for i, d := range docs {
|
||
c := domain.Citation{Index: i + 1, Content: d.Content}
|
||
if id, ok := d.MetaData["chunk_id"].(int64); ok {
|
||
c.ChunkId = id
|
||
}
|
||
if id, ok := d.MetaData["document_id"].(int64); ok {
|
||
c.DocumentId = id
|
||
}
|
||
if s, ok := d.MetaData["score"].(float64); ok {
|
||
c.Score = s
|
||
}
|
||
if srcs, ok := d.MetaData["sources"].([]string); ok {
|
||
c.Sources = srcs
|
||
}
|
||
cits = append(cits, c)
|
||
}
|
||
return cits
|
||
}
|
||
|
||
// buildSystemPrompt 系统提示词:引用资料编号 + 检索片段 + 知识图谱三元组(M5 图增强)
|
||
func buildSystemPrompt(citations []domain.Citation, triples []string) string {
|
||
var sb strings.Builder
|
||
sb.WriteString("你是一个本地知识库助手。请仅根据以下资料回答用户问题;若资料不足以回答,请明确说明。")
|
||
sb.WriteString("回答引用资料时,在对应位置标注 [编号]。\n\n【资料】\n")
|
||
for _, c := range citations {
|
||
sb.WriteString(fmt.Sprintf("[%d] %s\n", c.Index, c.Content))
|
||
}
|
||
if len(triples) > 0 {
|
||
sb.WriteString("\n【知识图谱】以下为与问题相关的实体关系,可辅助回答关系类问题:\n")
|
||
for _, t := range triples {
|
||
sb.WriteString(t + "\n")
|
||
}
|
||
}
|
||
return sb.String()
|
||
}
|
||
|
||
// ---------- HTTP 辅助 ----------
|
||
|
||
func buildOpenAIMessages(input []*schema.Message) []openAIMessage {
|
||
out := make([]openAIMessage, 0, len(input))
|
||
for _, m := range input {
|
||
if m == nil {
|
||
continue
|
||
}
|
||
role := string(m.Role)
|
||
if role == "" {
|
||
role = string(schema.User)
|
||
}
|
||
out = append(out, openAIMessage{Role: role, Content: m.Content})
|
||
}
|
||
return out
|
||
}
|
||
|
||
func postOpenAI(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) ([]byte, error) {
|
||
body, err := doOpenAIRequest(ctx, cfg, url, payload)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer body.Close()
|
||
resp, err := io.ReadAll(body)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return resp, nil
|
||
}
|
||
|
||
func postOpenAIStream(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) (io.ReadCloser, error) {
|
||
return doOpenAIRequest(ctx, cfg, url, payload)
|
||
}
|
||
|
||
func doOpenAIRequest(ctx context.Context, cfg *entity.ModelConfig, url string, payload any) (io.ReadCloser, error) {
|
||
start := time.Now()
|
||
buf, err := json.Marshal(payload)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(buf))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
req.Header.Set("Content-Type", "application/json")
|
||
if cfg.ApiKey != "" {
|
||
req.Header.Set("Authorization", "Bearer "+cfg.ApiKey)
|
||
}
|
||
resp, err := httpClient.Do(req)
|
||
if err != nil {
|
||
g.Log().Errorf(ctx, "model call failed: model=%s url=%s err=%v", cfg.ModelName, url, err)
|
||
return nil, err
|
||
}
|
||
if resp.StatusCode >= 400 {
|
||
msg, _ := io.ReadAll(resp.Body)
|
||
resp.Body.Close()
|
||
err := gerror.Newf("模型接口 %s 返回 %d: %s", url, resp.StatusCode, strings.TrimSpace(string(msg)))
|
||
g.Log().Errorf(ctx, "model call failed: model=%s url=%s status=%d err=%v", cfg.ModelName, url, resp.StatusCode, err)
|
||
return nil, err
|
||
}
|
||
g.Log().Infof(ctx, "model call ok: model=%s url=%s status=%d dur=%s", cfg.ModelName, url, resp.StatusCode, time.Since(start).Round(time.Millisecond))
|
||
return resp.Body, nil
|
||
}
|