132 lines
3.0 KiB
Go
132 lines
3.0 KiB
Go
package knowledge
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"fmt"
|
|
"github.com/google/uuid"
|
|
"resume-platform/internal/ai"
|
|
"resume-platform/internal/model"
|
|
"resume-platform/internal/repository"
|
|
)
|
|
|
|
type EmbeddingService struct {
|
|
aiProvider ai.Provider
|
|
repo repository.ResumeRepository
|
|
}
|
|
|
|
func NewEmbeddingService(provider ai.Provider, repo repository.ResumeRepository) *EmbeddingService {
|
|
return &EmbeddingService{aiProvider: provider, repo: repo}
|
|
}
|
|
|
|
func (s *EmbeddingService) Generate(ctx context.Context, text string) ([]float64, error) {
|
|
if text == "" {
|
|
return nil, fmt.Errorf("empty text")
|
|
}
|
|
|
|
contentHash := s.hashContent(text)
|
|
|
|
cachedEmbedding, err := s.repo.GetEmbeddingByContentHash(contentHash)
|
|
if err == nil && cachedEmbedding != nil && cachedEmbedding.Embedding != "" {
|
|
var embedding []float64
|
|
if err := json.Unmarshal([]byte(cachedEmbedding.Embedding), &embedding); err == nil && len(embedding) > 0 {
|
|
return embedding, nil
|
|
}
|
|
}
|
|
|
|
prompt := fmt.Sprintf(`请将以下文本转换为向量嵌入。直接输出JSON数组,不要包含任何其他内容。
|
|
|
|
文本内容:
|
|
%s`, text)
|
|
|
|
result, err := s.aiProvider.Generate(ctx, prompt)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
result = cleanEmbeddingResult(result)
|
|
|
|
var embedding []float64
|
|
if err := json.Unmarshal([]byte(result), &embedding); err != nil {
|
|
return nil, fmt.Errorf("failed to parse embedding: %w", err)
|
|
}
|
|
|
|
return embedding, nil
|
|
}
|
|
|
|
func (s *EmbeddingService) GenerateAndStore(ctx context.Context, text, sourceType, sourceID, sourceName, userID string) ([]float64, error) {
|
|
embedding, err := s.Generate(ctx, text)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
contentHash := s.hashContent(text)
|
|
embeddingJSON, _ := json.Marshal(embedding)
|
|
|
|
err = s.repo.CreateEmbedding(&model.Embedding{
|
|
ID: uuid.New().String(),
|
|
ContentHash: contentHash,
|
|
Embedding: string(embeddingJSON),
|
|
SourceType: sourceType,
|
|
SourceID: sourceID,
|
|
SourceName: sourceName,
|
|
UserID: userID,
|
|
})
|
|
|
|
return embedding, err
|
|
}
|
|
|
|
func (s *EmbeddingService) BatchGenerate(ctx context.Context, texts []string) ([][]float64, error) {
|
|
var results [][]float64
|
|
for _, text := range texts {
|
|
embedding, err := s.Generate(ctx, text)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
results = append(results, embedding)
|
|
}
|
|
return results, nil
|
|
}
|
|
|
|
func (s *EmbeddingService) hashContent(content string) string {
|
|
h := sha256.New()
|
|
h.Write([]byte(content))
|
|
return hex.EncodeToString(h.Sum(nil))
|
|
}
|
|
|
|
func cleanEmbeddingResult(result string) string {
|
|
result = trimString(result)
|
|
|
|
if idx := findFirstIndex(result, '['); idx >= 0 {
|
|
result = result[idx:]
|
|
}
|
|
if idx := findLastIndex(result, ']'); idx >= 0 {
|
|
result = result[:idx+1]
|
|
}
|
|
|
|
return result
|
|
}
|
|
|
|
func trimString(s string) string {
|
|
return s
|
|
}
|
|
|
|
func findFirstIndex(s string, c byte) int {
|
|
for i := 0; i < len(s); i++ {
|
|
if s[i] == c {
|
|
return i
|
|
}
|
|
}
|
|
return -1
|
|
}
|
|
|
|
func findLastIndex(s string, c byte) int {
|
|
for i := len(s) - 1; i >= 0; i-- {
|
|
if s[i] == c {
|
|
return i
|
|
}
|
|
}
|
|
return -1
|
|
} |