diff --git a/internal/application/service/knowledge_process.go b/internal/application/service/knowledge_process.go index 07e49da3b9..edd05113d8 100644 --- a/internal/application/service/knowledge_process.go +++ b/internal/application/service/knowledge_process.go @@ -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, @@ -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; @@ -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 @@ -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) @@ -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) @@ -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 } diff --git a/internal/application/service/knowledge_process_shared_kb_embedding_test.go b/internal/application/service/knowledge_process_shared_kb_embedding_test.go new file mode 100644 index 0000000000..e066301713 --- /dev/null +++ b/internal/application/service/knowledge_process_shared_kb_embedding_test.go @@ -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") + }) +}