Files
gochat/backend/internal/repository/article_embedding_repo.go
T
2026-08-14 22:40:48 +08:00

92 lines
3.1 KiB
Go

package repository
import (
"context"
"github.com/gochat/gochat/internal/model"
"github.com/pgvector/pgvector-go"
"gorm.io/gorm"
)
// ArticleEmbeddingRepo provides data access for ArticleEmbedding (semantic search).
type ArticleEmbeddingRepo struct {
db *gorm.DB
}
func NewArticleEmbeddingRepo(db *gorm.DB) *ArticleEmbeddingRepo {
return &ArticleEmbeddingRepo{db: db}
}
func (r *ArticleEmbeddingRepo) DB() *gorm.DB { return r.db }
// Upsert creates or updates an article embedding.
func (r *ArticleEmbeddingRepo) Upsert(ctx context.Context, emb *model.ArticleEmbedding) error {
// Try to find existing by article_id
var existing model.ArticleEmbedding
err := r.db.WithContext(ctx).Where("article_id = ?", emb.ArticleID).First(&existing).Error
if err == gorm.ErrRecordNotFound {
if r.db.Dialector != nil && r.db.Dialector.Name() == "sqlite" {
return r.db.WithContext(ctx).Omit("VectorEmbedding").Create(emb).Error
}
return r.db.WithContext(ctx).Create(emb).Error
}
if err != nil {
return err
}
// Update existing
existing.Embedding = emb.Embedding
existing.VectorEmbedding = emb.VectorEmbedding
existing.Term = emb.Term
if r.db.Dialector != nil && r.db.Dialector.Name() == "sqlite" {
return r.db.WithContext(ctx).Omit("VectorEmbedding").Save(&existing).Error
}
return r.db.WithContext(ctx).Save(&existing).Error
}
// GetByArticleID retrieves the embedding for a specific article.
func (r *ArticleEmbeddingRepo) GetByArticleID(ctx context.Context, articleID uint) (*model.ArticleEmbedding, error) {
var emb model.ArticleEmbedding
if err := r.db.WithContext(ctx).Where("article_id = ?", articleID).First(&emb).Error; err != nil {
return nil, err
}
return &emb, nil
}
// DeleteByArticleID removes the embedding for an article.
func (r *ArticleEmbeddingRepo) DeleteByArticleID(ctx context.Context, articleID uint) error {
return r.db.WithContext(ctx).Where("article_id = ?", articleID).Delete(&model.ArticleEmbedding{}).Error
}
// SearchByEmbedding performs semantic search using pgvector cosine distance.
// Returns articles sorted by similarity (closest first).
// Note: This method is PG-only (requires pgvector extension).
func (r *ArticleEmbeddingRepo) SearchByEmbedding(ctx context.Context, portalID uint, embedding pgvector.Vector, limit int) ([]model.Article, error) {
if limit <= 0 {
limit = 10
}
var rows []struct {
model.Article
Distance *float64 `gorm:"column:semantic_distance"`
}
// Join articles with article_embeddings, compute cosine distance on vector_embedding column.
// pgvector cosine distance operator: <=> (for vector type)
err := r.db.WithContext(ctx).Raw(`
SELECT a.*, ae.vector_embedding <=> ? AS semantic_distance FROM articles a
INNER JOIN article_embeddings ae ON ae.article_id = a.id
WHERE a.portal_id = ? AND a.status = 'published'
AND a.deleted_at IS NULL AND ae.deleted_at IS NULL
ORDER BY semantic_distance
LIMIT ?
`, embedding, portalID, limit).Scan(&rows).Error
if err != nil {
return nil, err
}
articles := make([]model.Article, len(rows))
for i := range rows {
articles[i] = rows[i].Article
articles[i].SemanticDistance = rows[i].Distance
}
return articles, nil
}