openai.go 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167
  1. package embedding
  2. import (
  3. "bytes"
  4. "context"
  5. "encoding/json"
  6. "fmt"
  7. "io"
  8. "log"
  9. "net/http"
  10. "strings"
  11. "time"
  12. )
  13. // OpenAIEmbeddingService OpenAI 嵌入服务实现
  14. type OpenAIEmbeddingService struct {
  15. apiURL string
  16. apiKey string
  17. model string
  18. dimension int
  19. }
  20. // NewOpenAIEmbeddingService 创建 OpenAI 嵌入服务实例
  21. func NewOpenAIEmbeddingService(apiURL, apiKey, model string) *OpenAIEmbeddingService {
  22. if apiURL == "" {
  23. apiURL = "https://api.openai.com/v1"
  24. }
  25. if model == "" {
  26. model = "text-embedding-3-small"
  27. }
  28. dimension := 1536 // text-embedding-3-small 的默认维度
  29. if model == "text-embedding-3-large" {
  30. dimension = 3072
  31. }
  32. return &OpenAIEmbeddingService{
  33. apiURL: apiURL,
  34. apiKey: apiKey,
  35. model: model,
  36. dimension: dimension,
  37. }
  38. }
  39. // EmbedText 向量化单个文本
  40. func (s *OpenAIEmbeddingService) EmbedText(ctx context.Context, text string) ([]float32, error) {
  41. vectors, err := s.EmbedTexts(ctx, []string{text})
  42. if err != nil {
  43. return nil, err
  44. }
  45. if len(vectors) == 0 {
  46. return nil, fmt.Errorf("未返回向量")
  47. }
  48. return vectors[0], nil
  49. }
  50. // EmbedTexts 批量向量化文本
  51. func (s *OpenAIEmbeddingService) EmbedTexts(ctx context.Context, texts []string) ([][]float32, error) {
  52. if len(texts) == 0 {
  53. return nil, nil
  54. }
  55. // 诊断日志:确认发请求前我们到底发了几条文本(用于排查 1 文档 vs 6 向量 问题)
  56. log.Printf("[嵌入] EmbedTexts 请求: len(texts)=%d, model=%s, apiURL=%s", len(texts), s.model, strings.TrimSuffix(s.apiURL, "/"))
  57. for i, t := range texts {
  58. runeLen := len([]rune(t))
  59. preview := t
  60. if runeLen > 60 {
  61. preview = string([]rune(t)[:60]) + "..."
  62. }
  63. log.Printf("[嵌入] texts[%d] 长度=%d 字符, 预览: %q", i, runeLen, preview)
  64. }
  65. // 支持填完整路径或仅填 base:若已以 /embeddings 结尾则不再追加,否则追加 /embeddings
  66. url := strings.TrimSuffix(s.apiURL, "/")
  67. if url != "" && !strings.HasSuffix(strings.ToLower(url), "/embeddings") {
  68. url = url + "/embeddings"
  69. } else if url == "" {
  70. url = s.apiURL + "/embeddings"
  71. }
  72. // 构建请求体
  73. requestBody := map[string]interface{}{
  74. "input": texts,
  75. "model": s.model,
  76. }
  77. jsonData, err := json.Marshal(requestBody)
  78. if err != nil {
  79. return nil, fmt.Errorf("序列化请求失败: %w", err)
  80. }
  81. // 创建 HTTP 请求
  82. req, err := http.NewRequestWithContext(ctx, "POST", url, bytes.NewBuffer(jsonData))
  83. if err != nil {
  84. return nil, fmt.Errorf("创建请求失败: %w", err)
  85. }
  86. req.Header.Set("Content-Type", "application/json")
  87. req.Header.Set("Authorization", "Bearer "+s.apiKey)
  88. // 发送请求
  89. client := &http.Client{Timeout: 30 * time.Second}
  90. resp, err := client.Do(req)
  91. if err != nil {
  92. return nil, fmt.Errorf("发送请求失败: %w", err)
  93. }
  94. defer resp.Body.Close()
  95. // 读取响应
  96. body, err := io.ReadAll(resp.Body)
  97. if err != nil {
  98. return nil, fmt.Errorf("读取响应失败: %w", err)
  99. }
  100. if resp.StatusCode != http.StatusOK {
  101. return nil, fmt.Errorf("OpenAI API 返回错误状态码 %d: %s", resp.StatusCode, string(body))
  102. }
  103. // 解析响应(若返回 HTML 则提示检查 API 地址/密钥)
  104. var response struct {
  105. Data []struct {
  106. Embedding []float64 `json:"embedding"`
  107. } `json:"data"`
  108. }
  109. if err := json.Unmarshal(body, &response); err != nil {
  110. if len(body) > 0 && body[0] == '<' {
  111. snippet := string(body)
  112. if len(snippet) > 200 {
  113. snippet = snippet[:200] + "..."
  114. }
  115. log.Printf("[嵌入] OpenAI 返回了 HTML 而非 JSON,请检查 API 地址与密钥。响应片段: %s", snippet)
  116. return nil, fmt.Errorf("嵌入 API 返回了 HTML 而非 JSON,请检查「设置 - 知识库向量模型」中的 API 地址与密钥: %w", err)
  117. }
  118. return nil, fmt.Errorf("解析响应失败: %w", err)
  119. }
  120. // 诊断日志:API 实际返回了几个向量(若与 len(texts) 不一致,说明接口按长文本分块或我们发了多条)
  121. numIn := len(texts)
  122. numOut := len(response.Data)
  123. log.Printf("[嵌入] EmbedTexts 响应: len(texts)=%d -> len(data)=%d (API 返回向量数)", numIn, numOut)
  124. if numOut != numIn {
  125. log.Printf("[嵌入] 数量不一致: 我们发了 %d 条文本,API 返回了 %d 个向量(可能接口对长文本分块,或请求被中间层改写)", numIn, numOut)
  126. }
  127. // 转换为 float32
  128. result := make([][]float32, len(response.Data))
  129. for i, item := range response.Data {
  130. result[i] = make([]float32, len(item.Embedding))
  131. for j, v := range item.Embedding {
  132. result[i][j] = float32(v)
  133. }
  134. }
  135. return result, nil
  136. }
  137. // GetDimension 获取向量维度
  138. func (s *OpenAIEmbeddingService) GetDimension() int {
  139. return s.dimension
  140. }
  141. // GetModelName 获取模型名称
  142. func (s *OpenAIEmbeddingService) GetModelName() string {
  143. return s.model
  144. }