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 { return r.db.WithContext(ctx).Create(emb).Error } if err != nil { return err } // Update existing existing.Embedding = emb.Embedding existing.Term = emb.Term 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 articles []model.Article // 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.* FROM articles a INNER JOIN article_embeddings ae ON ae.article_id = a.id WHERE a.portal_id = ? AND a.status = 'published' ORDER BY ae.vector_embedding <=> ? LIMIT ? `, portalID, embedding, limit).Scan(&articles).Error if err != nil { return nil, err } return articles, nil }