首次提交:初始化项目代码
This commit is contained in:
@@ -0,0 +1,132 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,362 @@
|
||||
package knowledge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"resume-platform/internal/model"
|
||||
"resume-platform/internal/repository"
|
||||
"resume-platform/pkg/logger"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type KnowledgeBaseService struct {
|
||||
embeddingService *EmbeddingService
|
||||
vectorService *VectorService
|
||||
repo repository.ResumeRepository
|
||||
}
|
||||
|
||||
func NewKnowledgeBaseService(embeddingService *EmbeddingService, vectorService *VectorService, repo repository.ResumeRepository) *KnowledgeBaseService {
|
||||
return &KnowledgeBaseService{
|
||||
embeddingService: embeddingService,
|
||||
vectorService: vectorService,
|
||||
repo: repo,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) AddDocument(ctx context.Context, documentID, content, userID string) error {
|
||||
if content == "" {
|
||||
return fmt.Errorf("empty content")
|
||||
}
|
||||
|
||||
chunks := s.chunkContent(content, 500, 50)
|
||||
|
||||
doc, _ := s.repo.GetDocumentByID(documentID)
|
||||
sourceName := documentID[:8]
|
||||
if doc != nil && doc.Name != "" {
|
||||
sourceName = doc.Name
|
||||
}
|
||||
|
||||
for i, chunk := range chunks {
|
||||
documentChunk := &model.DocumentChunk{
|
||||
ID: uuid.New().String(),
|
||||
DocumentID: documentID,
|
||||
ChunkIndex: i,
|
||||
Content: chunk,
|
||||
Metadata: "",
|
||||
}
|
||||
|
||||
err := s.vectorService.StoreChunk(ctx, documentChunk, userID, sourceName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
logger.Infof("Added %d chunks to document %s", len(chunks), documentID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) AddResume(ctx context.Context, resumeID string, resume *model.Resume) error {
|
||||
if resume == nil {
|
||||
return fmt.Errorf("nil resume")
|
||||
}
|
||||
|
||||
resumeText := s.serializeResume(resume)
|
||||
if resumeText == "" {
|
||||
return fmt.Errorf("empty resume content")
|
||||
}
|
||||
|
||||
chunks := s.chunkContent(resumeText, 500, 50)
|
||||
sourceName := resume.BasicInfo.Name
|
||||
if sourceName == "" {
|
||||
sourceName = resumeID[:8]
|
||||
}
|
||||
|
||||
for i, chunk := range chunks {
|
||||
documentChunk := &model.DocumentChunk{
|
||||
ID: uuid.New().String(),
|
||||
DocumentID: resumeID,
|
||||
ChunkIndex: i,
|
||||
Content: chunk,
|
||||
Metadata: "",
|
||||
}
|
||||
|
||||
err := s.vectorService.StoreChunk(ctx, documentChunk, resume.UserID, sourceName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
err := s.vectorService.Store(ctx, resumeText, "resume", resumeID, sourceName, resume.UserID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
logger.Infof("Added resume %s to knowledge base with %d chunks", resumeID, len(chunks))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) AddText(ctx context.Context, text, sourceType, sourceID, sourceName, userID string) error {
|
||||
if text == "" {
|
||||
return fmt.Errorf("empty text")
|
||||
}
|
||||
|
||||
return s.vectorService.Store(ctx, text, sourceType, sourceID, sourceName, userID)
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) Retrieve(ctx context.Context, query string, topK int, userID string) ([]model.ChunkWithScore, error) {
|
||||
return s.vectorService.SearchWithChunks(ctx, query, topK, userID)
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) RetrieveAll(ctx context.Context, query string, topK int) ([]model.ChunkWithScore, error) {
|
||||
return s.vectorService.Search(ctx, query, topK, "")
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) DeleteDocument(ctx context.Context, documentID string) error {
|
||||
err := s.repo.DeleteEmbeddingsBySource("document", documentID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.vectorService.DeleteDocumentChunks(ctx, documentID)
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) DeleteResume(ctx context.Context, resumeID string) error {
|
||||
return s.repo.DeleteEmbeddingsBySource("resume", resumeID)
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) GetStats(ctx context.Context, userID string) (map[string]int, error) {
|
||||
var stats = map[string]int{}
|
||||
|
||||
if userID != "" {
|
||||
embeddings, err := s.repo.GetEmbeddingsByUserID(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stats["embeddings"] = len(embeddings)
|
||||
|
||||
chunks, err := s.repo.GetDocumentChunksByUserID(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stats["chunks"] = len(chunks)
|
||||
} else {
|
||||
embeddings, err := s.repo.GetAllEmbeddings()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stats["embeddings"] = len(embeddings)
|
||||
|
||||
chunks, err := s.repo.GetAllDocumentChunks()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
stats["chunks"] = len(chunks)
|
||||
}
|
||||
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) chunkContent(content string, chunkSize, overlap int) []string {
|
||||
if content == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
content = strings.ReplaceAll(content, "\r\n", "\n")
|
||||
content = strings.ReplaceAll(content, "\r", "\n")
|
||||
|
||||
var chunks []string
|
||||
start := 0
|
||||
contentLen := len(content)
|
||||
|
||||
for start < contentLen {
|
||||
end := start + chunkSize
|
||||
if end > contentLen {
|
||||
end = contentLen
|
||||
}
|
||||
|
||||
if end < contentLen {
|
||||
for i := end; i > start && i > start+chunkSize-overlap; i-- {
|
||||
c := rune(content[i])
|
||||
if c == '\n' || c == ';' || c == '\r' {
|
||||
end = i + 1
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
chunk := strings.TrimSpace(content[start:end])
|
||||
if chunk != "" {
|
||||
chunks = append(chunks, chunk)
|
||||
}
|
||||
|
||||
start = end - overlap
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
if start >= contentLen {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return chunks
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) serializeResume(resume *model.Resume) string {
|
||||
var builder strings.Builder
|
||||
|
||||
if resume.BasicInfo.Name != "" {
|
||||
builder.WriteString("姓名:")
|
||||
builder.WriteString(resume.BasicInfo.Name)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.Title != "" {
|
||||
builder.WriteString("职位:")
|
||||
builder.WriteString(resume.BasicInfo.Title)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.Email != "" {
|
||||
builder.WriteString("邮箱:")
|
||||
builder.WriteString(resume.BasicInfo.Email)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.Phone != "" {
|
||||
builder.WriteString("电话:")
|
||||
builder.WriteString(resume.BasicInfo.Phone)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.Location != "" {
|
||||
builder.WriteString("所在地:")
|
||||
builder.WriteString(resume.BasicInfo.Location)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.Summary != "" {
|
||||
builder.WriteString("个人简介:")
|
||||
builder.WriteString(resume.BasicInfo.Summary)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if resume.BasicInfo.JobTarget != "" {
|
||||
builder.WriteString("求职目标:")
|
||||
builder.WriteString(resume.BasicInfo.JobTarget)
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
|
||||
if len(resume.Experience) > 0 {
|
||||
builder.WriteString("\n【工作经历】\n")
|
||||
for _, exp := range resume.Experience {
|
||||
builder.WriteString("公司:")
|
||||
builder.WriteString(exp.Company)
|
||||
builder.WriteString("\n职位:")
|
||||
builder.WriteString(exp.Position)
|
||||
builder.WriteString("\n时间:")
|
||||
builder.WriteString(exp.StartDate)
|
||||
if exp.EndDate != "" {
|
||||
builder.WriteString(" - ")
|
||||
builder.WriteString(exp.EndDate)
|
||||
}
|
||||
builder.WriteString("\n职责:")
|
||||
builder.WriteString(exp.Description)
|
||||
builder.WriteString("\n")
|
||||
if len(exp.Highlights) > 0 {
|
||||
builder.WriteString("亮点:")
|
||||
builder.WriteString(strings.Join(exp.Highlights, ";"))
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
if len(resume.Education) > 0 {
|
||||
builder.WriteString("\n【教育背景】\n")
|
||||
for _, edu := range resume.Education {
|
||||
builder.WriteString("学校:")
|
||||
builder.WriteString(edu.School)
|
||||
builder.WriteString("\n学位:")
|
||||
builder.WriteString(edu.Degree)
|
||||
builder.WriteString("\n专业:")
|
||||
builder.WriteString(edu.Major)
|
||||
builder.WriteString("\n时间:")
|
||||
builder.WriteString(edu.StartDate)
|
||||
if edu.EndDate != "" {
|
||||
builder.WriteString(" - ")
|
||||
builder.WriteString(edu.EndDate)
|
||||
}
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
if len(resume.Skills) > 0 {
|
||||
builder.WriteString("\n【专业技能】\n")
|
||||
for _, skill := range resume.Skills {
|
||||
builder.WriteString(skill.Name)
|
||||
if skill.Level != "" {
|
||||
builder.WriteString("(")
|
||||
builder.WriteString(skill.Level)
|
||||
builder.WriteString(")")
|
||||
}
|
||||
if skill.Category != "" {
|
||||
builder.WriteString(" - ")
|
||||
builder.WriteString(skill.Category)
|
||||
}
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
if len(resume.Projects) > 0 {
|
||||
builder.WriteString("\n【项目经验】\n")
|
||||
for _, proj := range resume.Projects {
|
||||
builder.WriteString("项目名称:")
|
||||
builder.WriteString(proj.Name)
|
||||
builder.WriteString("\n描述:")
|
||||
builder.WriteString(proj.Description)
|
||||
builder.WriteString("\n")
|
||||
if len(proj.TechStack) > 0 {
|
||||
builder.WriteString("技术栈:")
|
||||
builder.WriteString(strings.Join(proj.TechStack, "、"))
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if len(proj.Highlights) > 0 {
|
||||
builder.WriteString("亮点:")
|
||||
builder.WriteString(strings.Join(proj.Highlights, ";"))
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
if len(proj.Achievements) > 0 {
|
||||
builder.WriteString("成果:")
|
||||
builder.WriteString(strings.Join(proj.Achievements, ";"))
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
builder.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
result := builder.String()
|
||||
if len(result) > 15000 {
|
||||
result = result[:15000]
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *KnowledgeBaseService) BuildIndex(ctx context.Context) error {
|
||||
documents, err := s.repo.GetAllDocumentChunks()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, chunk := range documents {
|
||||
if chunk.Embedding == "" && chunk.Content != "" {
|
||||
embedding, err := s.embeddingService.Generate(ctx, chunk.Content)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
embeddingJSON, _ := json.Marshal(embedding)
|
||||
chunk.Embedding = string(embeddingJSON)
|
||||
err = s.repo.CreateDocumentChunk(chunk)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package knowledge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"resume-platform/internal/ai"
|
||||
"resume-platform/internal/prompt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type RAGService struct {
|
||||
aiProvider ai.Provider
|
||||
knowledgeBaseService *KnowledgeBaseService
|
||||
}
|
||||
|
||||
func NewRAGService(provider ai.Provider, knowledgeBaseService *KnowledgeBaseService) *RAGService {
|
||||
return &RAGService{aiProvider: provider, knowledgeBaseService: knowledgeBaseService}
|
||||
}
|
||||
|
||||
func (s *RAGService) Generate(ctx context.Context, query string, userID string) (string, error) {
|
||||
contextResults, err := s.knowledgeBaseService.Retrieve(ctx, query, 5, userID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var contextStr string
|
||||
for _, result := range contextResults {
|
||||
contextStr += fmt.Sprintf("【来源:%s】\n%s\n\n", result.SourceName, result.Content)
|
||||
}
|
||||
|
||||
if contextStr == "" {
|
||||
return s.aiProvider.Generate(ctx, query)
|
||||
}
|
||||
|
||||
promptText := prompt.BuildRAGAnswerPrompt(contextStr, query)
|
||||
|
||||
return s.aiProvider.Generate(ctx, promptText)
|
||||
}
|
||||
|
||||
func (s *RAGService) GenerateWithContext(ctx context.Context, query string, context []string) (string, error) {
|
||||
contextStr := strings.Join(context, "\n\n")
|
||||
|
||||
promptText := prompt.BuildRAGAnswerPrompt(contextStr, query)
|
||||
|
||||
return s.aiProvider.Generate(ctx, promptText)
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
package knowledge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"resume-platform/internal/model"
|
||||
"resume-platform/internal/repository"
|
||||
"sort"
|
||||
)
|
||||
|
||||
type VectorService struct {
|
||||
embeddingService *EmbeddingService
|
||||
repo repository.ResumeRepository
|
||||
}
|
||||
|
||||
func NewVectorService(embeddingService *EmbeddingService, repo repository.ResumeRepository) *VectorService {
|
||||
return &VectorService{embeddingService: embeddingService, repo: repo}
|
||||
}
|
||||
|
||||
func (s *VectorService) Store(ctx context.Context, content, sourceType, sourceID, sourceName, userID string) error {
|
||||
_, err := s.embeddingService.GenerateAndStore(ctx, content, sourceType, sourceID, sourceName, userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *VectorService) StoreChunk(ctx context.Context, chunk *model.DocumentChunk, userID string, sourceName string) error {
|
||||
if chunk.Content == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
embedding, err := s.embeddingService.Generate(ctx, chunk.Content)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
embeddingJSON, _ := json.Marshal(embedding)
|
||||
chunk.Embedding = string(embeddingJSON)
|
||||
|
||||
err = s.repo.CreateDocumentChunk(chunk)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
contentHash := s.embeddingService.hashContent(chunk.Content)
|
||||
err = s.repo.CreateEmbedding(&model.Embedding{
|
||||
ID: chunk.ID,
|
||||
ContentHash: contentHash,
|
||||
Embedding: string(embeddingJSON),
|
||||
SourceType: "document",
|
||||
SourceID: chunk.DocumentID,
|
||||
SourceName: sourceName,
|
||||
UserID: userID,
|
||||
})
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *VectorService) Search(ctx context.Context, query string, topK int, userID string) ([]model.ChunkWithScore, error) {
|
||||
if query == "" {
|
||||
return nil, fmt.Errorf("empty query")
|
||||
}
|
||||
|
||||
queryEmbedding, err := s.embeddingService.Generate(ctx, query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var allEmbeddings []*model.Embedding
|
||||
if userID != "" {
|
||||
allEmbeddings, err = s.repo.GetEmbeddingsByUserID(userID)
|
||||
} else {
|
||||
allEmbeddings, err = s.repo.GetAllEmbeddings()
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var results []model.ChunkWithScore
|
||||
for _, emb := range allEmbeddings {
|
||||
if emb.Embedding == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
var embedding []float64
|
||||
if err := json.Unmarshal([]byte(emb.Embedding), &embedding); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
score := cosineSimilarity(queryEmbedding, embedding)
|
||||
if score > 0.3 {
|
||||
results = append(results, model.ChunkWithScore{
|
||||
Content: emb.SourceName,
|
||||
Score: score,
|
||||
SourceID: emb.SourceID,
|
||||
SourceName: emb.SourceName,
|
||||
SourceType: emb.SourceType,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(results, func(i, j int) bool {
|
||||
return results[i].Score > results[j].Score
|
||||
})
|
||||
|
||||
if len(results) > topK {
|
||||
results = results[:topK]
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *VectorService) SearchWithChunks(ctx context.Context, query string, topK int, userID string) ([]model.ChunkWithScore, error) {
|
||||
if query == "" {
|
||||
return nil, fmt.Errorf("empty query")
|
||||
}
|
||||
|
||||
queryEmbedding, err := s.embeddingService.Generate(ctx, query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var allChunks []*model.DocumentChunk
|
||||
if userID != "" {
|
||||
allChunks, err = s.repo.GetDocumentChunksByUserID(userID)
|
||||
} else {
|
||||
allChunks, err = s.repo.GetAllDocumentChunks()
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var results []model.ChunkWithScore
|
||||
for _, chunk := range allChunks {
|
||||
if chunk.Embedding == "" || chunk.Content == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
var embedding []float64
|
||||
if err := json.Unmarshal([]byte(chunk.Embedding), &embedding); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
score := cosineSimilarity(queryEmbedding, embedding)
|
||||
if score > 0.3 {
|
||||
doc, _ := s.repo.GetDocumentByID(chunk.DocumentID)
|
||||
sourceName := ""
|
||||
if doc != nil {
|
||||
sourceName = doc.Name
|
||||
}
|
||||
results = append(results, model.ChunkWithScore{
|
||||
Content: chunk.Content,
|
||||
Score: score,
|
||||
SourceID: chunk.DocumentID,
|
||||
SourceName: sourceName,
|
||||
SourceType: "document",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(results, func(i, j int) bool {
|
||||
return results[i].Score > results[j].Score
|
||||
})
|
||||
|
||||
if len(results) > topK {
|
||||
results = results[:topK]
|
||||
}
|
||||
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func (s *VectorService) Delete(ctx context.Context, sourceType, sourceID string) error {
|
||||
return s.repo.DeleteEmbeddingsBySource(sourceType, sourceID)
|
||||
}
|
||||
|
||||
func (s *VectorService) DeleteDocumentChunks(ctx context.Context, documentID string) error {
|
||||
_, err := s.repo.GetDocumentChunksByDocumentID(documentID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func cosineSimilarity(a, b []float64) float64 {
|
||||
if len(a) != len(b) {
|
||||
return 0
|
||||
}
|
||||
|
||||
var dotProduct, magA, magB float64
|
||||
for i := range a {
|
||||
dotProduct += a[i] * b[i]
|
||||
magA += a[i] * a[i]
|
||||
magB += b[i] * b[i]
|
||||
}
|
||||
|
||||
if magA == 0 || magB == 0 {
|
||||
return 0
|
||||
}
|
||||
|
||||
return dotProduct / (math.Sqrt(magA) * math.Sqrt(magB))
|
||||
}
|
||||
Reference in New Issue
Block a user