This repository has been archived on 2026-04-03. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
core/server/store/postgres/knowledge_bases.go
2026-02-27 02:38:35 +00:00

450 lines
15 KiB
Go

package postgres
import (
"context"
"database/sql"
"fmt"
"strings"
"git.gobha.me/xcaliber/chat-switchboard/models"
"github.com/lib/pq"
)
// ── KnowledgeBaseStore ──────────────────────────
type KnowledgeBaseStore struct{}
func NewKnowledgeBaseStore() *KnowledgeBaseStore { return &KnowledgeBaseStore{} }
// ── KB CRUD ──────────────────────────────────────
func (s *KnowledgeBaseStore) Create(ctx context.Context, kb *models.KnowledgeBase) error {
return DB.QueryRowContext(ctx, `
INSERT INTO knowledge_bases (name, description, scope, owner_id, team_id, embedding_config, status)
VALUES ($1, $2, $3, $4, $5, $6, $7)
RETURNING id, document_count, chunk_count, total_bytes, created_at, updated_at`,
kb.Name, kb.Description, kb.Scope,
models.NullString(kb.OwnerID), models.NullString(kb.TeamID),
ToJSON(kb.EmbeddingConfig), kb.Status,
).Scan(&kb.ID, &kb.DocumentCount, &kb.ChunkCount, &kb.TotalBytes, &kb.CreatedAt, &kb.UpdatedAt)
}
func (s *KnowledgeBaseStore) GetByID(ctx context.Context, id string) (*models.KnowledgeBase, error) {
kb, err := scanKB(DB.QueryRowContext(ctx, `
SELECT id, name, description, scope, owner_id, team_id, embedding_config,
document_count, chunk_count, total_bytes, status, created_at, updated_at
FROM knowledge_bases WHERE id = $1`, id))
if err != nil {
return nil, err
}
return kb, nil
}
func (s *KnowledgeBaseStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
b := NewUpdate("knowledge_bases")
for k, v := range fields {
if k == "embedding_config" {
b.SetJSON(k, v)
} else {
b.Set(k, v)
}
}
b.Set("updated_at", "now()")
if !b.HasSets() {
return nil
}
b.Where("id", id)
_, err := b.Exec(DB)
return err
}
func (s *KnowledgeBaseStore) Delete(ctx context.Context, id string) error {
// CASCADE deletes kb_documents, kb_chunks, channel_knowledge_bases.
_, err := DB.ExecContext(ctx, "DELETE FROM knowledge_bases WHERE id = $1", id)
return err
}
// ── Scoped Listing ──────────────────────────────
func (s *KnowledgeBaseStore) ListForUser(ctx context.Context, userID string, teamIDs []string) ([]models.KnowledgeBase, error) {
// User can see: global + their teams + personal + group-granted.
q := `
SELECT id, name, description, scope, owner_id, team_id, embedding_config,
document_count, chunk_count, total_bytes, status, created_at, updated_at
FROM knowledge_bases
WHERE scope = 'global'
OR (scope = 'personal' AND owner_id = $1)`
args := []interface{}{userID}
if len(teamIDs) > 0 {
q += fmt.Sprintf(` OR (scope = 'team' AND team_id = ANY($%d))`, len(args)+1)
args = append(args, pq.Array(teamIDs))
}
// Group-granted KBs
q += fmt.Sprintf(`
OR id IN (
SELECT rg.resource_id FROM resource_grants rg
WHERE rg.resource_type = 'knowledge_base'
AND (
rg.grant_scope = 'global'
OR (rg.grant_scope = 'groups'
AND EXISTS (
SELECT 1 FROM group_members gm
WHERE gm.user_id = $%d
AND gm.group_id = ANY(rg.granted_groups)
))
)
)`, len(args)+1)
args = append(args, userID)
q += ` ORDER BY name`
return queryKBs(ctx, q, args...)
}
func (s *KnowledgeBaseStore) ListGlobal(ctx context.Context) ([]models.KnowledgeBase, error) {
return queryKBs(ctx, `
SELECT id, name, description, scope, owner_id, team_id, embedding_config,
document_count, chunk_count, total_bytes, status, created_at, updated_at
FROM knowledge_bases WHERE scope = 'global' ORDER BY name`)
}
func (s *KnowledgeBaseStore) ListForTeam(ctx context.Context, teamID string) ([]models.KnowledgeBase, error) {
return queryKBs(ctx, `
SELECT id, name, description, scope, owner_id, team_id, embedding_config,
document_count, chunk_count, total_bytes, status, created_at, updated_at
FROM knowledge_bases WHERE team_id = $1 ORDER BY name`, teamID)
}
func (s *KnowledgeBaseStore) ListPersonal(ctx context.Context, userID string) ([]models.KnowledgeBase, error) {
return queryKBs(ctx, `
SELECT id, name, description, scope, owner_id, team_id, embedding_config,
document_count, chunk_count, total_bytes, status, created_at, updated_at
FROM knowledge_bases WHERE scope = 'personal' AND owner_id = $1 ORDER BY name`, userID)
}
// ── Documents ────────────────────────────────────
func (s *KnowledgeBaseStore) CreateDocument(ctx context.Context, doc *models.KBDocument) error {
return DB.QueryRowContext(ctx, `
INSERT INTO kb_documents (kb_id, filename, content_type, size_bytes, storage_key, status, uploaded_by)
VALUES ($1, $2, $3, $4, $5, $6, $7)
RETURNING id, chunk_count, created_at, updated_at`,
doc.KBID, doc.Filename, doc.ContentType, doc.SizeBytes,
doc.StorageKey, doc.Status, doc.UploadedBy,
).Scan(&doc.ID, &doc.ChunkCount, &doc.CreatedAt, &doc.UpdatedAt)
}
func (s *KnowledgeBaseStore) GetDocument(ctx context.Context, id string) (*models.KBDocument, error) {
var doc models.KBDocument
var extractedText, errMsg sql.NullString
err := DB.QueryRowContext(ctx, `
SELECT id, kb_id, filename, content_type, size_bytes, storage_key,
extracted_text, chunk_count, status, error, uploaded_by, created_at, updated_at
FROM kb_documents WHERE id = $1`, id).Scan(
&doc.ID, &doc.KBID, &doc.Filename, &doc.ContentType, &doc.SizeBytes,
&doc.StorageKey, &extractedText, &doc.ChunkCount, &doc.Status,
&errMsg, &doc.UploadedBy, &doc.CreatedAt, &doc.UpdatedAt,
)
if err != nil {
return nil, err
}
doc.ExtractedText = NullableStringPtr(extractedText)
doc.Error = NullableStringPtr(errMsg)
return &doc, nil
}
func (s *KnowledgeBaseStore) ListDocuments(ctx context.Context, kbID string) ([]models.KBDocument, error) {
rows, err := DB.QueryContext(ctx, `
SELECT id, kb_id, filename, content_type, size_bytes, storage_key,
chunk_count, status, error, uploaded_by, created_at, updated_at
FROM kb_documents WHERE kb_id = $1 ORDER BY created_at`, kbID)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.KBDocument
for rows.Next() {
var doc models.KBDocument
var errMsg sql.NullString
err := rows.Scan(&doc.ID, &doc.KBID, &doc.Filename, &doc.ContentType,
&doc.SizeBytes, &doc.StorageKey, &doc.ChunkCount, &doc.Status,
&errMsg, &doc.UploadedBy, &doc.CreatedAt, &doc.UpdatedAt)
if err != nil {
return nil, err
}
doc.Error = NullableStringPtr(errMsg)
result = append(result, doc)
}
return result, rows.Err()
}
func (s *KnowledgeBaseStore) UpdateDocumentStatus(ctx context.Context, id string, status string, errMsg *string) error {
_, err := DB.ExecContext(ctx, `
UPDATE kb_documents SET status = $2, error = $3, updated_at = now()
WHERE id = $1`, id, status, models.NullString(errMsg))
return err
}
func (s *KnowledgeBaseStore) UpdateDocumentText(ctx context.Context, id string, text string, chunkCount int) error {
_, err := DB.ExecContext(ctx, `
UPDATE kb_documents SET extracted_text = $2, chunk_count = $3, updated_at = now()
WHERE id = $1`, id, text, chunkCount)
return err
}
func (s *KnowledgeBaseStore) DeleteDocument(ctx context.Context, id string) (*models.KBDocument, error) {
var doc models.KBDocument
var errMsg sql.NullString
// Return the row before deleting so caller can clean up storage.
err := DB.QueryRowContext(ctx, `
DELETE FROM kb_documents WHERE id = $1
RETURNING id, kb_id, filename, storage_key, status, error`,
id).Scan(&doc.ID, &doc.KBID, &doc.Filename, &doc.StorageKey, &doc.Status, &errMsg)
if err != nil {
return nil, err
}
doc.Error = NullableStringPtr(errMsg)
return &doc, nil
}
// ── Chunks ───────────────────────────────────────
func (s *KnowledgeBaseStore) InsertChunks(ctx context.Context, chunks []models.KBChunk) error {
if len(chunks) == 0 {
return nil
}
// Batch insert using a single multi-row INSERT.
// Embedding vectors are inserted as text representations that pgvector parses.
const cols = 7 // kb_id, document_id, chunk_index, content, token_count, embedding, metadata
valueParts := make([]string, 0, len(chunks))
args := make([]interface{}, 0, len(chunks)*cols)
for i, c := range chunks {
base := i * cols
valueParts = append(valueParts, fmt.Sprintf(
"($%d, $%d, $%d, $%d, $%d, $%d::vector, $%d)",
base+1, base+2, base+3, base+4, base+5, base+6, base+7,
))
args = append(args, c.KBID, c.DocumentID, c.ChunkIndex,
c.Content, c.TokenCount, vectorToString(c.Embedding), ToJSON(c.Metadata))
}
q := fmt.Sprintf(`INSERT INTO kb_chunks (kb_id, document_id, chunk_index, content, token_count, embedding, metadata)
VALUES %s`, strings.Join(valueParts, ", "))
_, err := DB.ExecContext(ctx, q, args...)
return err
}
func (s *KnowledgeBaseStore) DeleteChunksForDocument(ctx context.Context, documentID string) error {
_, err := DB.ExecContext(ctx, "DELETE FROM kb_chunks WHERE document_id = $1", documentID)
return err
}
func (s *KnowledgeBaseStore) SimilaritySearch(ctx context.Context, kbIDs []string, queryVec []float64, threshold float64, limit int) ([]models.KBSearchResult, error) {
if len(kbIDs) == 0 || len(queryVec) == 0 {
return nil, nil
}
if limit <= 0 {
limit = 5
}
rows, err := DB.QueryContext(ctx, `
SELECT c.content, c.metadata, d.filename, kb.name,
1 - (c.embedding <=> $1::vector) AS similarity
FROM kb_chunks c
JOIN kb_documents d ON c.document_id = d.id
JOIN knowledge_bases kb ON c.kb_id = kb.id
WHERE c.kb_id = ANY($2)
AND 1 - (c.embedding <=> $1::vector) > $3
ORDER BY c.embedding <=> $1::vector
LIMIT $4`,
vectorToString(queryVec), pq.Array(kbIDs), threshold, limit,
)
if err != nil {
return nil, err
}
defer rows.Close()
var results []models.KBSearchResult
for rows.Next() {
var r models.KBSearchResult
var metadataJSON []byte
err := rows.Scan(&r.Content, &metadataJSON, &r.Filename, &r.KBName, &r.Similarity)
if err != nil {
return nil, err
}
ScanJSON(metadataJSON, &r.Metadata)
results = append(results, r)
}
return results, rows.Err()
}
// ── Channel Links ─────────────────────────────────
func (s *KnowledgeBaseStore) SetChannelKBs(ctx context.Context, channelID string, kbIDs []string) error {
tx, err := DB.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
// Clear existing links.
_, err = tx.ExecContext(ctx, "DELETE FROM channel_knowledge_bases WHERE channel_id = $1", channelID)
if err != nil {
return err
}
// Insert new links.
for _, kbID := range kbIDs {
_, err = tx.ExecContext(ctx, `
INSERT INTO channel_knowledge_bases (channel_id, kb_id, enabled) VALUES ($1, $2, true)`,
channelID, kbID)
if err != nil {
return err
}
}
return tx.Commit()
}
func (s *KnowledgeBaseStore) GetChannelKBs(ctx context.Context, channelID string) ([]models.ChannelKB, error) {
rows, err := DB.QueryContext(ctx, `
SELECT ckb.kb_id, kb.name, ckb.enabled, kb.document_count
FROM channel_knowledge_bases ckb
JOIN knowledge_bases kb ON ckb.kb_id = kb.id
WHERE ckb.channel_id = $1
ORDER BY kb.name`, channelID)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.ChannelKB
for rows.Next() {
var ckb models.ChannelKB
if err := rows.Scan(&ckb.KBID, &ckb.KBName, &ckb.Enabled, &ckb.DocumentCount); err != nil {
return nil, err
}
result = append(result, ckb)
}
return result, rows.Err()
}
func (s *KnowledgeBaseStore) GetActiveKBIDs(ctx context.Context, channelID string, userID string, teamIDs []string) ([]string, error) {
// Return KB IDs that are: enabled on this channel AND user has access to.
q := `
SELECT ckb.kb_id
FROM channel_knowledge_bases ckb
JOIN knowledge_bases kb ON ckb.kb_id = kb.id
WHERE ckb.channel_id = $1 AND ckb.enabled = true
AND (
kb.scope = 'global'
OR (kb.scope = 'personal' AND kb.owner_id = $2)`
args := []interface{}{channelID, userID}
if len(teamIDs) > 0 {
q += fmt.Sprintf(` OR (kb.scope = 'team' AND kb.team_id = ANY($%d))`, len(args)+1)
args = append(args, pq.Array(teamIDs))
}
q += `)`
rows, err := DB.QueryContext(ctx, q, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var ids []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
return ids, rows.Err()
}
// ── Stats ────────────────────────────────────────
func (s *KnowledgeBaseStore) UpdateStats(ctx context.Context, kbID string) error {
_, err := DB.ExecContext(ctx, `
UPDATE knowledge_bases SET
document_count = (SELECT COUNT(*) FROM kb_documents WHERE kb_id = $1 AND status != 'error'),
chunk_count = (SELECT COUNT(*) FROM kb_chunks WHERE kb_id = $1),
total_bytes = COALESCE((SELECT SUM(size_bytes) FROM kb_documents WHERE kb_id = $1), 0),
updated_at = now()
WHERE id = $1`, kbID)
return err
}
// ── Scan Helpers ─────────────────────────────────
func scanKB(row *sql.Row) (*models.KnowledgeBase, error) {
var kb models.KnowledgeBase
var ownerID, teamID sql.NullString
var embCfgJSON []byte
err := row.Scan(
&kb.ID, &kb.Name, &kb.Description, &kb.Scope,
&ownerID, &teamID, &embCfgJSON,
&kb.DocumentCount, &kb.ChunkCount, &kb.TotalBytes,
&kb.Status, &kb.CreatedAt, &kb.UpdatedAt,
)
if err != nil {
return nil, err
}
kb.OwnerID = NullableStringPtr(ownerID)
kb.TeamID = NullableStringPtr(teamID)
ScanJSON(embCfgJSON, &kb.EmbeddingConfig)
return &kb, nil
}
func queryKBs(ctx context.Context, query string, args ...interface{}) ([]models.KnowledgeBase, error) {
rows, err := DB.QueryContext(ctx, query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var result []models.KnowledgeBase
for rows.Next() {
var kb models.KnowledgeBase
var ownerID, teamID sql.NullString
var embCfgJSON []byte
err := rows.Scan(
&kb.ID, &kb.Name, &kb.Description, &kb.Scope,
&ownerID, &teamID, &embCfgJSON,
&kb.DocumentCount, &kb.ChunkCount, &kb.TotalBytes,
&kb.Status, &kb.CreatedAt, &kb.UpdatedAt,
)
if err != nil {
return nil, err
}
kb.OwnerID = NullableStringPtr(ownerID)
kb.TeamID = NullableStringPtr(teamID)
ScanJSON(embCfgJSON, &kb.EmbeddingConfig)
result = append(result, kb)
}
return result, rows.Err()
}
// vectorToString converts a float64 slice to pgvector text format: [0.1,0.2,0.3]
func vectorToString(v []float64) string {
if len(v) == 0 {
return ""
}
parts := make([]string, len(v))
for i, f := range v {
parts[i] = fmt.Sprintf("%g", f)
}
return "[" + strings.Join(parts, ",") + "]"
}