Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 26 additions & 5 deletions internal/application/service/knowledge_process.go
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,27 @@ func buildParentChildConfigs(cc types.ChunkingConfig, base chunker.SplitterConfi
return chunker.DeriveParentChildConfigs(base, cc.ParentChunkSize, cc.ChildChunkSize)
}

// resolveKBEmbeddingModel resolves a knowledge base's embedding model under the
// tenant that owns it.
//
// A knowledge base shared into another tenant is written by the owner's tenant
// but processed with the viewer's tenant in ctx, so the plain ctx-tenant lookup
// searches a tenant that never had the model row and fails with
// "Model not found" — aborting document processing for every document in the
// shared knowledge base. Resolving under kb.TenantID mirrors what the search
// path already does (see knowledgeBaseService.GetQueryEmbedding), so a shared
// knowledge base embeds and queries with the same provider and vector space.
//
// A ctx with no tenant falls through to the plain lookup rather than panicking
// here: the missing tenant is the model service's error to report, and it did
// so before this branch existed.
func (s *knowledgeService) resolveKBEmbeddingModel(ctx context.Context, kb *types.KnowledgeBase) (embedding.Embedder, error) {
if currentTenantID, ok := types.TenantIDFromContext(ctx); ok && kb.TenantID != currentTenantID {
return s.modelService.GetEmbeddingModelForTenant(ctx, kb.EmbeddingModelID, kb.TenantID)
}
return s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
}

// processChunks processes chunks and creates embeddings for knowledge content
func (s *knowledgeService) processChunks(ctx context.Context,
kb *types.KnowledgeBase, knowledge *types.Knowledge, chunks []types.ParsedChunk,
Expand Down Expand Up @@ -388,7 +409,7 @@ func (s *knowledgeService) processChunks(ctx context.Context,
var embeddingModel embedding.Embedder
if kb.NeedsEmbeddingModel() {
var err error
embeddingModel, err = s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
embeddingModel, err = s.resolveKBEmbeddingModel(ctx, kb)
if err != nil {
// Terminal for this attempt, and it has to be recorded as such.
// A KB that indexes vectors cannot proceed without an embedder;
Expand Down Expand Up @@ -1586,7 +1607,7 @@ func (s *knowledgeService) ProcessSummaryGeneration(ctx context.Context, t *asyn
return fmt.Errorf("failed to init retrieve engine: %w", err)
}

embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
embeddingModel, err := s.resolveKBEmbeddingModel(ctx, kb)
if err != nil {
logger.Errorf(ctx, "Failed to get embedding model: %v", err)
summaryErr = err
Expand Down Expand Up @@ -1856,7 +1877,7 @@ func (s *knowledgeService) processQuestionGenerationForKnowledge(ctx context.Con
resolvedModelID = kb.SummaryModelID

// Initialize embedding model and retrieval engine
embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
embeddingModel, err := s.resolveKBEmbeddingModel(ctx, kb)
if err != nil {
exitStatus = "get_embedding_model_failed"
logger.Errorf(ctx, "Failed to get embedding model: %v", err)
Expand Down Expand Up @@ -2166,7 +2187,7 @@ func (s *knowledgeService) processQuestionGenerationForChunks(ctx context.Contex
}
resolvedModelID = kb.SummaryModelID

embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, kb.EmbeddingModelID)
embeddingModel, err := s.resolveKBEmbeddingModel(ctx, kb)
if err != nil {
exitStatus = "get_embedding_model_failed"
logger.Errorf(ctx, "Failed to get embedding model: %v", err)
Expand Down Expand Up @@ -3113,7 +3134,7 @@ func (s *knowledgeService) updateChunkVector(ctx context.Context, kbID string, c
if !sourceKB.NeedsEmbeddingModel() {
return nil
}
embeddingModel, err := s.modelService.GetEmbeddingModel(ctx, sourceKB.EmbeddingModelID)
embeddingModel, err := s.resolveKBEmbeddingModel(ctx, sourceKB)
if err != nil {
return err
}
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,217 @@
package service

import (
"context"
"errors"
"testing"

"github.com/Tencent/WeKnora/internal/models/embedding"
"github.com/Tencent/WeKnora/internal/types"
"github.com/Tencent/WeKnora/internal/types/interfaces"
"github.com/stretchr/testify/require"
)

// sharedKBModelService stands in for a deployment where each tenant owns its own
// model rows. ownerTenantID owns embeddingModelID, so the plain ctx-tenant
// lookup reports "Model not found" for anyone else — which is exactly the
// failure a shared knowledge base used to hit, and which the two counters make
// observable.
type sharedKBModelService struct {
interfaces.ModelService
embedder embedding.Embedder
ownerTenantID uint64
embeddingModel string
ctxLookups []uint64
ownerLookups []uint64
ownerLookupErrs error
}

func (s *sharedKBModelService) GetEmbeddingModel(
ctx context.Context, modelID string,
) (embedding.Embedder, error) {
tenantID, _ := types.TenantIDFromContext(ctx)
s.ctxLookups = append(s.ctxLookups, tenantID)
if tenantID != s.ownerTenantID {
return nil, errors.New("Model not found")
}
return s.embedder, nil
}

func (s *sharedKBModelService) GetEmbeddingModelForTenant(
_ context.Context, _ string, tenantID uint64,
) (embedding.Embedder, error) {
s.ownerLookups = append(s.ownerLookups, tenantID)
if tenantID != s.ownerTenantID {
return nil, errors.New("Model not found")
}
if s.ownerLookupErrs != nil {
return nil, s.ownerLookupErrs
}
return s.embedder, nil
}

// sharedKBService wires the same collaborators as the other processChunks
// tests, with a model service that only knows the owning tenant's row.
func sharedKBService(
knowledge *types.Knowledge, modelSvc *sharedKBModelService,
) (*knowledgeService, *parentChildChunkService, *parentChildRetrieveEngine) {
chunkService := &parentChildChunkService{}
retrieveEngine := &parentChildRetrieveEngine{}
svc := &knowledgeService{
repo: &parentChildKnowledgeRepo{knowledge: knowledge},
chunkRepo: chunkService,
modelService: modelSvc,
retrieveEngine: parentChildRetrieveRegistry{engine: retrieveEngine},
graphEngine: parentChildGraphRepo{},
tenantRepo: parentChildTenantRepo{},
task: parentChildTaskEnqueuer{},
}
return svc, chunkService, retrieveEngine
}

// sharedKBContext is the request context a viewer of a shared knowledge base
// carries: the owner's tenant in the knowledge row, the viewer's in ctx. The
// tenant object the vector store is resolved from belongs to the viewer too —
// the model row is the only thing that lives with the owner.
func sharedKBContext(viewerTenant uint64) context.Context {
tenant := &types.Tenant{
ID: viewerTenant,
RetrieverEngines: types.RetrieverEngines{Engines: []types.RetrieverEngineParams{
{
RetrieverType: types.VectorRetrieverType,
RetrieverEngineType: types.PostgresRetrieverEngineType,
},
}},
}
ctx := context.WithValue(context.Background(), types.TenantIDContextKey, viewerTenant)
return context.WithValue(ctx, types.TenantInfoContextKey, tenant)
}

// A knowledge base shared into a shared space is written by the owning tenant
// but parsed with the viewer's tenant in ctx. The embedding model belongs to
// the owner, so resolving it by ctx alone reported "Model not found" and
// processChunks failed the document — leaving every document in the shared
// knowledge base unsearchable for anyone but its owner.
func TestProcessChunksIndexesSharedKnowledgeBaseUnderOwningTenant(t *testing.T) {
const (
ownerTenant uint64 = 10000
viewerTenant uint64 = 10001
)
knowledge := &types.Knowledge{
ID: "knowledge-1",
TenantID: ownerTenant,
KnowledgeBaseID: "kb-1",
ParseStatus: types.ParseStatusProcessing,
}
modelSvc := &sharedKBModelService{
embedder: parentChildEmbedder{},
ownerTenantID: ownerTenant,
embeddingModel: "embedding-1",
}
svc, chunkService, retrieveEngine := sharedKBService(knowledge, modelSvc)
kb := &types.KnowledgeBase{
ID: "kb-1",
TenantID: ownerTenant,
EmbeddingModelID: "embedding-1",
IndexingStrategy: types.IndexingStrategy{VectorEnabled: true},
}

svc.processChunks(sharedKBContext(viewerTenant), kb, knowledge,
[]types.ParsedChunk{{Content: "shared body", Seq: 0, Start: 0, End: 11}})

require.Equal(t, []uint64{ownerTenant}, modelSvc.ownerLookups,
"a shared knowledge base must resolve its embedding model under the owning tenant")
require.Empty(t, modelSvc.ctxLookups,
"the viewer's tenant does not own the model row, so the ctx lookup must not be reached")
require.Len(t, chunkService.created, 1, "the document must be chunked, not failed")
require.Len(t, retrieveEngine.indexed, 1, "and indexed, so the viewer can search it")
}

// The owning tenant processing its own knowledge base is the ordinary path and
// must keep using the plain ctx-tenant lookup, so nothing changes where there
// is no cross-tenant sharing involved.
func TestProcessChunksKeepsCtxLookupForOwnKnowledgeBase(t *testing.T) {
const ownerTenant uint64 = 10000
knowledge := &types.Knowledge{
ID: "knowledge-1",
TenantID: ownerTenant,
KnowledgeBaseID: "kb-1",
ParseStatus: types.ParseStatusProcessing,
}
modelSvc := &sharedKBModelService{
embedder: parentChildEmbedder{},
ownerTenantID: ownerTenant,
embeddingModel: "embedding-1",
}
svc, chunkService, _ := sharedKBService(knowledge, modelSvc)
kb := &types.KnowledgeBase{
ID: "kb-1",
TenantID: ownerTenant,
EmbeddingModelID: "embedding-1",
IndexingStrategy: types.IndexingStrategy{VectorEnabled: true},
}

svc.processChunks(sharedKBContext(ownerTenant), kb, knowledge,
[]types.ParsedChunk{{Content: "owned body", Seq: 0, Start: 0, End: 11}})

require.Equal(t, []uint64{ownerTenant}, modelSvc.ctxLookups,
"an owned knowledge base still resolves through the ctx tenant")
require.Empty(t, modelSvc.ownerLookups,
"the cross-tenant branch is for shared knowledge bases only")
require.Len(t, chunkService.created, 1)
}

// The helper is shared by every processing stage, so pin the branch itself and
// not only its effect through processChunks: a provider failure on the owner
// side must still surface, and a ctx without a tenant must not panic here —
// the model service reports that on its own terms, as it did before.
func TestResolveKBEmbeddingModel(t *testing.T) {
const (
ownerTenant uint64 = 10000
viewerTenant uint64 = 10001
)
shared := &types.KnowledgeBase{
ID: "kb-1", TenantID: ownerTenant, EmbeddingModelID: "embedding-1",
}
owned := &types.KnowledgeBase{
ID: "kb-2", TenantID: ownerTenant, EmbeddingModelID: "embedding-1",
}
newSvc := func() (*knowledgeService, *sharedKBModelService) {
modelSvc := &sharedKBModelService{
embedder: parentChildEmbedder{}, ownerTenantID: ownerTenant, embeddingModel: "embedding-1",
}
return &knowledgeService{modelService: modelSvc}, modelSvc
}

t.Run("shared knowledge base resolves under the owner", func(t *testing.T) {
svc, modelSvc := newSvc()
embedder, err := svc.resolveKBEmbeddingModel(sharedKBContext(viewerTenant), shared)
require.NoError(t, err)
require.NotNil(t, embedder)
require.Equal(t, []uint64{ownerTenant}, modelSvc.ownerLookups)
require.Empty(t, modelSvc.ctxLookups)
})

t.Run("owned knowledge base resolves through ctx", func(t *testing.T) {
svc, modelSvc := newSvc()
embedder, err := svc.resolveKBEmbeddingModel(sharedKBContext(ownerTenant), owned)
require.NoError(t, err)
require.NotNil(t, embedder)
require.Equal(t, []uint64{ownerTenant}, modelSvc.ctxLookups)
require.Empty(t, modelSvc.ownerLookups)
})

t.Run("a missing tenant is the model service's error to report", func(t *testing.T) {
svc, modelSvc := newSvc()
_, err := svc.resolveKBEmbeddingModel(context.Background(), shared)
require.Error(t, err)
require.Empty(t, modelSvc.ownerLookups)
})

t.Run("an owner-side provider failure still surfaces", func(t *testing.T) {
svc, modelSvc := newSvc()
modelSvc.ownerLookupErrs = errors.New("dial tcp: connection refused")
_, err := svc.resolveKBEmbeddingModel(sharedKBContext(viewerTenant), shared)
require.ErrorContains(t, err, "connection refused")
})
}