mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-09-24 16:29:01 +08:00
refactor: Improve Elasticsearch query construction and error handling
This commit is contained in:
@@ -353,38 +353,70 @@ func (e *elasticsearchRepository) deleteByFieldList(ctx context.Context, field s
|
||||
// Returns a JSON string representing the query conditions
|
||||
func (e *elasticsearchRepository) getBaseConds(params typesLocal.RetrieveParams) string {
|
||||
// Build MUST conditions (positive filters)
|
||||
must := make([]string, 0)
|
||||
must := make([]map[string]interface{}, 0)
|
||||
if len(params.KnowledgeBaseIDs) > 0 {
|
||||
ids, _ := json.Marshal(params.KnowledgeBaseIDs)
|
||||
must = append(must, fmt.Sprintf(`{"terms": {"knowledge_base_id.keyword": %s}}`, ids))
|
||||
must = append(must, map[string]interface{}{
|
||||
"terms": map[string]interface{}{
|
||||
"knowledge_base_id.keyword": params.KnowledgeBaseIDs,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Build MUST_NOT conditions (negative filters)
|
||||
mustNot := make([]string, 0)
|
||||
mustNot := make([]map[string]interface{}, 0)
|
||||
// Exclude disabled chunks (is_enabled = false)
|
||||
// Note: Historical data without is_enabled field will be included (not matching must_not)
|
||||
mustNot = append(mustNot, `{"term": {"is_enabled": false}}`)
|
||||
mustNot = append(mustNot, map[string]interface{}{
|
||||
"term": map[string]interface{}{
|
||||
"is_enabled": false,
|
||||
},
|
||||
})
|
||||
if len(params.ExcludeKnowledgeIDs) > 0 {
|
||||
ids, _ := json.Marshal(params.ExcludeKnowledgeIDs)
|
||||
mustNot = append(mustNot, fmt.Sprintf(`{"terms": {"knowledge_id.keyword": %s}}`, ids))
|
||||
mustNot = append(mustNot, map[string]interface{}{
|
||||
"terms": map[string]interface{}{
|
||||
"knowledge_id.keyword": params.ExcludeKnowledgeIDs,
|
||||
},
|
||||
})
|
||||
}
|
||||
if len(params.ExcludeChunkIDs) > 0 {
|
||||
ids, _ := json.Marshal(params.ExcludeChunkIDs)
|
||||
mustNot = append(mustNot, fmt.Sprintf(`{"terms": {"chunk_id.keyword": %s}}`, ids))
|
||||
mustNot = append(mustNot, map[string]interface{}{
|
||||
"terms": map[string]interface{}{
|
||||
"chunk_id.keyword": params.ExcludeChunkIDs,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// Combine conditions based on presence
|
||||
var query map[string]interface{}
|
||||
if len(must) == 0 && len(mustNot) == 0 {
|
||||
return "{}" // Empty query if no conditions
|
||||
query = map[string]interface{}{}
|
||||
} else if len(must) == 0 {
|
||||
query = map[string]interface{}{
|
||||
"bool": map[string]interface{}{
|
||||
"must_not": mustNot,
|
||||
},
|
||||
}
|
||||
} else if len(mustNot) == 0 {
|
||||
query = map[string]interface{}{
|
||||
"bool": map[string]interface{}{
|
||||
"must": must,
|
||||
},
|
||||
}
|
||||
} else {
|
||||
query = map[string]interface{}{
|
||||
"bool": map[string]interface{}{
|
||||
"must": must,
|
||||
"must_not": mustNot,
|
||||
},
|
||||
}
|
||||
}
|
||||
if len(must) == 0 {
|
||||
return fmt.Sprintf(`{"bool": {"must_not": [%s]}}`, strings.Join(mustNot, ","))
|
||||
|
||||
// Marshal to JSON string
|
||||
jsonBytes, err := json.Marshal(query)
|
||||
if err != nil {
|
||||
return "{}"
|
||||
}
|
||||
if len(mustNot) == 0 {
|
||||
return fmt.Sprintf(`{"bool": {"must": [%s]}}`, strings.Join(must, ","))
|
||||
}
|
||||
return fmt.Sprintf(`{"bool": {"must": [%s], "must_not": [%s]}}`,
|
||||
strings.Join(must, ","), strings.Join(mustNot, ","))
|
||||
return string(jsonBytes)
|
||||
}
|
||||
|
||||
func (e *elasticsearchRepository) Retrieve(ctx context.Context,
|
||||
@@ -439,26 +471,43 @@ func (e *elasticsearchRepository) buildVectorSearchQuery(ctx context.Context,
|
||||
) (string, error) {
|
||||
log := logger.GetLogger(ctx)
|
||||
|
||||
filter := e.getBaseConds(params)
|
||||
|
||||
// Serialize the query vector
|
||||
queryVectorJSON, err := json.Marshal(params.Embedding)
|
||||
if err != nil {
|
||||
log.Errorf("[ElasticsearchV7] Failed to marshal query vector: %v", err)
|
||||
return "", fmt.Errorf("failed to marshal query embedding: %w", err)
|
||||
// Parse filter conditions
|
||||
var filterQuery map[string]interface{}
|
||||
filterJSON := e.getBaseConds(params)
|
||||
if err := json.Unmarshal([]byte(filterJSON), &filterQuery); err != nil {
|
||||
log.Errorf("[ElasticsearchV7] Failed to unmarshal filter: %v", err)
|
||||
filterQuery = map[string]interface{}{}
|
||||
}
|
||||
|
||||
// Construct the script_score query
|
||||
query := fmt.Sprintf(
|
||||
`{"query":{"script_score":{"query":{"bool":{"filter":[%s]}},
|
||||
"script":{"source":"cosineSimilarity(params.query_vector,'embedding')",
|
||||
"params":{"query_vector":%s}},"min_score":%f}},"size":%d}`,
|
||||
filter,
|
||||
string(queryVectorJSON),
|
||||
params.Threshold,
|
||||
params.TopK,
|
||||
)
|
||||
// Construct the script_score query using structured objects
|
||||
queryObj := map[string]interface{}{
|
||||
"query": map[string]interface{}{
|
||||
"script_score": map[string]interface{}{
|
||||
"query": map[string]interface{}{
|
||||
"bool": map[string]interface{}{
|
||||
"filter": []interface{}{filterQuery},
|
||||
},
|
||||
},
|
||||
"script": map[string]interface{}{
|
||||
"source": "cosineSimilarity(params.query_vector,'embedding')",
|
||||
"params": map[string]interface{}{
|
||||
"query_vector": params.Embedding,
|
||||
},
|
||||
},
|
||||
"min_score": params.Threshold,
|
||||
},
|
||||
},
|
||||
"size": params.TopK,
|
||||
}
|
||||
|
||||
// Marshal to JSON string
|
||||
queryBytes, err := json.Marshal(queryObj)
|
||||
if err != nil {
|
||||
log.Errorf("[ElasticsearchV7] Failed to marshal query: %v", err)
|
||||
return "", fmt.Errorf("failed to marshal query: %w", err)
|
||||
}
|
||||
|
||||
query := string(queryBytes)
|
||||
log.Debugf("[ElasticsearchV7] Executing vector search with query: %s", query)
|
||||
return query, nil
|
||||
}
|
||||
|
||||
@@ -145,6 +145,9 @@ func (b *graphBuilder) extractEntities(ctx context.Context, chunk *types.Chunk)
|
||||
defer b.mutex.Unlock()
|
||||
|
||||
for _, entity := range extractedEntities {
|
||||
if entity == nil {
|
||||
continue
|
||||
}
|
||||
if entity.Title == "" || entity.Description == "" {
|
||||
log.WithField("entity", entity).Warn("Invalid entity with empty title or description")
|
||||
continue
|
||||
@@ -272,6 +275,10 @@ func (b *graphBuilder) extractRelationships(ctx context.Context,
|
||||
relationship.Source, relationship.Target, relationship.ID)
|
||||
} else {
|
||||
// This relationship already exists, update its properties
|
||||
if existingRel == nil {
|
||||
log.Warnf("existingRel is nil, skip update")
|
||||
continue
|
||||
}
|
||||
chunkIDsAdded := 0
|
||||
for _, chunkID := range relationChunkIDs {
|
||||
if !slices.Contains(existingRel.ChunkIDs, chunkID) {
|
||||
|
||||
@@ -433,7 +433,6 @@ func (s *knowledgeBaseService) HybridSearch(ctx context.Context,
|
||||
logger.Infof(ctx, "Hybrid search parameters, knowledge base ID: %s, query text: %s", id, params.QueryText)
|
||||
|
||||
tenantInfo := ctx.Value(types.TenantInfoContextKey).(*types.Tenant)
|
||||
logger.Infof(ctx, "Creating composite retrieval engine, tenant ID: %d", tenantInfo.ID)
|
||||
|
||||
// Create a composite retrieval engine with tenant's configured retrievers
|
||||
retrieveEngine, err := retriever.NewCompositeRetrieveEngine(s.retrieveEngine, tenantInfo.RetrieverEngines.Engines)
|
||||
|
||||
@@ -44,7 +44,6 @@ func (s *mcpServiceService) CreateMCPService(ctx context.Context, service *types
|
||||
return fmt.Errorf("failed to create MCP service: %w", err)
|
||||
}
|
||||
|
||||
logger.GetLogger(ctx).Infof("MCP service created: %s (ID: %s)", service.Name, service.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -587,15 +587,13 @@ func (s *sessionService) selectChatModelIDWithOverride(ctx context.Context, sess
|
||||
// 3. Session's SummaryModelID (if not Remote)
|
||||
// 4. First knowledge base's SummaryModelID
|
||||
func (s *sessionService) selectChatModelID(ctx context.Context, session *types.Session, knowledgeBaseIDs []string) (string, error) {
|
||||
chatModelID := ""
|
||||
|
||||
// First, check if session has a SummaryModelID and if it's a Remote model
|
||||
if session.SummaryModelID != "" {
|
||||
model, err := s.modelService.GetModelByID(ctx, session.SummaryModelID)
|
||||
if err == nil && model != nil && model.Source == types.ModelSourceRemote {
|
||||
chatModelID = session.SummaryModelID
|
||||
logger.Infof(ctx, "Using session's Remote summary model: %s", chatModelID)
|
||||
return chatModelID, nil
|
||||
logger.Infof(ctx, "Using session's Remote summary model: %s", session.SummaryModelID)
|
||||
return session.SummaryModelID, nil
|
||||
} else if err == nil && model != nil {
|
||||
// Session has a model but it's not Remote, we'll check knowledge bases for Remote models
|
||||
logger.Infof(ctx, "Session has summary model %s but it's not Remote, checking knowledge bases for Remote models", session.SummaryModelID)
|
||||
@@ -603,7 +601,7 @@ func (s *sessionService) selectChatModelID(ctx context.Context, session *types.S
|
||||
}
|
||||
|
||||
// If no Remote model found from session, check knowledge bases for Remote models
|
||||
if chatModelID == "" && len(knowledgeBaseIDs) > 0 {
|
||||
if len(knowledgeBaseIDs) > 0 {
|
||||
// Try to find a knowledge base with Remote model
|
||||
for _, kbID := range knowledgeBaseIDs {
|
||||
kb, err := s.knowledgeBaseService.GetKnowledgeBaseByID(ctx, kbID)
|
||||
@@ -614,44 +612,35 @@ func (s *sessionService) selectChatModelID(ctx context.Context, session *types.S
|
||||
if kb != nil && kb.SummaryModelID != "" {
|
||||
model, err := s.modelService.GetModelByID(ctx, kb.SummaryModelID)
|
||||
if err == nil && model != nil && model.Source == types.ModelSourceRemote {
|
||||
chatModelID = kb.SummaryModelID
|
||||
logger.Infof(ctx, "Using Remote summary model from knowledge base %s: %s", kbID, chatModelID)
|
||||
return chatModelID, nil
|
||||
logger.Infof(ctx, "Using Remote summary model from knowledge base %s: %s", kbID, kb.SummaryModelID)
|
||||
return kb.SummaryModelID, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If still no Remote model found, use session's SummaryModelID if available
|
||||
if chatModelID == "" && session.SummaryModelID != "" {
|
||||
chatModelID = session.SummaryModelID
|
||||
logger.Infof(ctx, "No Remote model found, using session's summary model: %s", chatModelID)
|
||||
return chatModelID, nil
|
||||
if session.SummaryModelID != "" {
|
||||
logger.Infof(ctx, "No Remote model found, using session's summary model: %s", session.SummaryModelID)
|
||||
return session.SummaryModelID, nil
|
||||
}
|
||||
|
||||
// If still empty, use first knowledge base's model
|
||||
if chatModelID == "" {
|
||||
kb, err := s.knowledgeBaseService.GetKnowledgeBaseByID(ctx, knowledgeBaseIDs[0])
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to get knowledge base for model ID: %v", err)
|
||||
return "", fmt.Errorf("failed to get knowledge base %s: %w", knowledgeBaseIDs[0], err)
|
||||
}
|
||||
if kb != nil && kb.SummaryModelID != "" {
|
||||
chatModelID = kb.SummaryModelID
|
||||
logger.Infof(ctx, "Using summary model from first knowledge base %s: %s", knowledgeBaseIDs[0], chatModelID)
|
||||
return chatModelID, nil
|
||||
} else {
|
||||
logger.Errorf(ctx, "Knowledge base %s has no summary model ID", knowledgeBaseIDs[0])
|
||||
return "", fmt.Errorf("knowledge base %s has no summary model configured", knowledgeBaseIDs[0])
|
||||
}
|
||||
kb, err := s.knowledgeBaseService.GetKnowledgeBaseByID(ctx, knowledgeBaseIDs[0])
|
||||
if err != nil {
|
||||
logger.Errorf(ctx, "Failed to get knowledge base for model ID: %v", err)
|
||||
return "", fmt.Errorf("failed to get knowledge base %s: %w", knowledgeBaseIDs[0], err)
|
||||
}
|
||||
if kb != nil && kb.SummaryModelID != "" {
|
||||
logger.Infof(ctx, "Using summary model from first knowledge base %s: %s", knowledgeBaseIDs[0], kb.SummaryModelID)
|
||||
return kb.SummaryModelID, nil
|
||||
} else {
|
||||
logger.Errorf(ctx, "Knowledge base %s has no summary model ID", knowledgeBaseIDs[0])
|
||||
return "", fmt.Errorf("knowledge base %s has no summary model configured", knowledgeBaseIDs[0])
|
||||
}
|
||||
}
|
||||
|
||||
if chatModelID == "" {
|
||||
logger.Error(ctx, "No chat model ID available")
|
||||
return "", errors.New("no chat model ID available: session has no SummaryModelID and knowledge bases have no SummaryModelID")
|
||||
}
|
||||
|
||||
return chatModelID, nil
|
||||
logger.Error(ctx, "No chat model ID available")
|
||||
return "", errors.New("no chat model ID available: session has no SummaryModelID and knowledge bases have no SummaryModelID")
|
||||
}
|
||||
|
||||
// KnowledgeQAByEvent processes knowledge QA through a series of events in the pipeline
|
||||
|
||||
@@ -18,7 +18,14 @@ import (
|
||||
)
|
||||
|
||||
// JWT secret key - in production this should be from environment variable
|
||||
var jwtSecret = []byte("your-secret-key")
|
||||
|
||||
func getJwtSecret() string {
|
||||
jwtSecret := os.Getenv("JWT_SECRET")
|
||||
if jwtSecret != "" {
|
||||
return jwtSecret
|
||||
}
|
||||
return "your-secret-key"
|
||||
}
|
||||
|
||||
// userService implements the UserService interface
|
||||
type userService struct {
|
||||
@@ -282,7 +289,7 @@ func (s *userService) GenerateTokens(ctx context.Context, user *types.User) (acc
|
||||
}
|
||||
|
||||
accessTokenObj := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims)
|
||||
accessToken, err = accessTokenObj.SignedString(jwtSecret)
|
||||
accessToken, err = accessTokenObj.SignedString([]byte(getJwtSecret()))
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
@@ -296,7 +303,7 @@ func (s *userService) GenerateTokens(ctx context.Context, user *types.User) (acc
|
||||
}
|
||||
|
||||
refreshTokenObj := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims)
|
||||
refreshToken, err = refreshTokenObj.SignedString(jwtSecret)
|
||||
refreshToken, err = refreshTokenObj.SignedString([]byte(getJwtSecret()))
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
@@ -334,7 +341,7 @@ func (s *userService) ValidateToken(ctx context.Context, tokenString string) (*t
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
||||
}
|
||||
return jwtSecret, nil
|
||||
return []byte(getJwtSecret()), nil
|
||||
})
|
||||
|
||||
if err != nil || !token.Valid {
|
||||
@@ -366,7 +373,7 @@ func (s *userService) RefreshToken(ctx context.Context, refreshTokenString strin
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
||||
}
|
||||
return jwtSecret, nil
|
||||
return []byte(getJwtSecret()), nil
|
||||
})
|
||||
|
||||
if err != nil || !token.Valid {
|
||||
|
||||
@@ -73,7 +73,6 @@ func (h *ChunkHandler) GetChunkByIDOnly(c *gin.Context) {
|
||||
chunk.Content = secutils.SanitizeForDisplay(chunk.Content)
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Successfully retrieved chunk by ID, chunk ID: %s", chunkID)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": chunk,
|
||||
@@ -102,9 +101,6 @@ func (h *ChunkHandler) ListKnowledgeChunks(c *gin.Context) {
|
||||
|
||||
chunkType := []types.ChunkType{types.ChunkTypeText}
|
||||
|
||||
logger.Infof(ctx, "Retrieving knowledge chunks list, knowledge ID: %s, page: %d, page size: %d",
|
||||
knowledgeID, pagination.Page, pagination.PageSize)
|
||||
|
||||
// Use pagination for query
|
||||
result, err := h.service.ListPagedChunksByKnowledgeID(ctx, knowledgeID, &pagination, chunkType)
|
||||
if err != nil {
|
||||
@@ -120,10 +116,6 @@ func (h *ChunkHandler) ListKnowledgeChunks(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
logger.Infof(
|
||||
ctx, "Successfully retrieved knowledge chunks list, knowledge ID: %s, total: %d",
|
||||
knowledgeID, result.Total,
|
||||
)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": result.Data,
|
||||
@@ -221,8 +213,6 @@ func (h *ChunkHandler) UpdateChunk(c *gin.Context) {
|
||||
|
||||
chunk.IsEnabled = req.IsEnabled
|
||||
|
||||
logger.Infof(ctx, "Updating knowledge chunk, knowledge ID: %s, chunk ID: %s", knowledgeID, chunk.ID)
|
||||
|
||||
if err := h.service.UpdateChunk(ctx, chunk); err != nil {
|
||||
logger.ErrorWithFields(ctx, err, nil)
|
||||
c.Error(errors.NewInternalServerError(err.Error()))
|
||||
@@ -242,21 +232,18 @@ func (h *ChunkHandler) DeleteChunk(c *gin.Context) {
|
||||
logger.Info(ctx, "Start deleting knowledge chunk")
|
||||
|
||||
// Validate parameters and get chunk
|
||||
chunk, knowledgeID, err := h.validateAndGetChunk(c)
|
||||
chunk, _, err := h.validateAndGetChunk(c)
|
||||
if err != nil {
|
||||
c.Error(err)
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Deleting knowledge chunk, knowledge ID: %s, chunk ID: %s", knowledgeID, chunk.ID)
|
||||
|
||||
if err := h.service.DeleteChunk(ctx, chunk.ID); err != nil {
|
||||
logger.ErrorWithFields(ctx, err, nil)
|
||||
c.Error(errors.NewInternalServerError(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Knowledge chunk deleted successfully, knowledge ID: %s, chunk ID: %s", knowledgeID, chunk.ID)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "Chunk deleted",
|
||||
@@ -275,16 +262,6 @@ func (h *ChunkHandler) DeleteChunksByKnowledgeID(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// Get tenant ID from context
|
||||
tenantID, exists := c.Get(types.TenantIDContextKey.String())
|
||||
if !exists {
|
||||
logger.Error(ctx, "Failed to get tenant ID")
|
||||
c.Error(errors.NewUnauthorizedError("Unauthorized"))
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Deleting all chunks under knowledge, knowledge ID: %s, tenant ID: %d", knowledgeID, tenantID.(uint))
|
||||
|
||||
// Delete all chunks under the knowledge
|
||||
err := h.service.DeleteChunksByKnowledgeID(ctx, knowledgeID)
|
||||
if err != nil {
|
||||
@@ -293,7 +270,6 @@ func (h *ChunkHandler) DeleteChunksByKnowledgeID(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "All chunks under knowledge deleted successfully, knowledge ID: %s", knowledgeID)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "All chunks under knowledge deleted",
|
||||
|
||||
@@ -51,6 +51,12 @@ func (e *EvaluationHandler) Evaluation(c *gin.Context) {
|
||||
logger.Infof(ctx, "Executing evaluation, tenant: %v, dataset: %s, knowledge_base: %s, chat: %s, rerank: %s",
|
||||
tenantID, request.DatasetID, request.KnowledgeBaseID, request.ChatModelID, request.RerankModelID)
|
||||
|
||||
if request == nil {
|
||||
logger.Error(ctx, "Request is nil")
|
||||
c.Error(errors.NewBadRequestError("Invalid request parameters"))
|
||||
return
|
||||
}
|
||||
|
||||
task, err := e.evaluationService.Evaluation(ctx,
|
||||
request.DatasetID,
|
||||
request.KnowledgeBaseID,
|
||||
@@ -81,12 +87,13 @@ func (e *EvaluationHandler) GetEvaluationResult(c *gin.Context) {
|
||||
|
||||
logger.Info(ctx, "Start retrieving evaluation result")
|
||||
|
||||
var request *GetEvaluationRequest
|
||||
var request GetEvaluationRequest
|
||||
if err := c.ShouldBind(&request); err != nil {
|
||||
logger.Error(ctx, "Failed to parse request parameters", err)
|
||||
c.Error(errors.NewBadRequestError("Invalid request parameters").WithDetails(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
result, err := e.evaluationService.EvaluationResult(ctx, request.TaskID)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, nil)
|
||||
|
||||
@@ -362,7 +362,7 @@ func (h *InitializationHandler) InitializeByKB(c *gin.Context) {
|
||||
}
|
||||
|
||||
if kb == nil {
|
||||
logger.Error(ctx, "Knowledge base not found", "kbId", kbIdStr)
|
||||
logger.Error(ctx, "Knowledge base not found")
|
||||
c.Error(errors.NewNotFoundError("知识库不存在"))
|
||||
return
|
||||
}
|
||||
@@ -584,6 +584,9 @@ func (h *InitializationHandler) InitializeByKB(c *gin.Context) {
|
||||
// 找到模型ID
|
||||
var embeddingModelID, llmModelID, vlmModelID string
|
||||
for _, model := range processedModels {
|
||||
if model == nil {
|
||||
continue
|
||||
}
|
||||
if model.Type == types.ModelTypeEmbedding {
|
||||
embeddingModelID = model.ID
|
||||
}
|
||||
@@ -671,7 +674,6 @@ func (h *InitializationHandler) InitializeByKB(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Info(ctx, "Knowledge base configuration updated successfully", "kbId", kbIdStr)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "知识库配置更新成功",
|
||||
@@ -806,15 +808,11 @@ func (h *InitializationHandler) DownloadOllamaModel(c *gin.Context) {
|
||||
// 检查模型是否已存在
|
||||
available, err := h.ollamaService.IsModelAvailable(ctx, req.ModelName)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"model_name": req.ModelName,
|
||||
})
|
||||
c.Error(errors.NewInternalServerError("检查模型状态失败: " + err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if available {
|
||||
logger.Infof(ctx, "Model %s already exists", req.ModelName)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "模型已存在",
|
||||
@@ -869,7 +867,7 @@ func (h *InitializationHandler) DownloadOllamaModel(c *gin.Context) {
|
||||
h.downloadModelAsync(newCtx, taskID, req.ModelName)
|
||||
}()
|
||||
|
||||
logger.Infof(ctx, "Created download task for model: %s, task ID: %s", req.ModelName, taskID)
|
||||
logger.Infof(ctx, "Created download task for model, task ID: %s", taskID)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"message": "模型下载任务已创建",
|
||||
@@ -956,7 +954,7 @@ func (h *InitializationHandler) ListOllamaModels(c *gin.Context) {
|
||||
func (h *InitializationHandler) downloadModelAsync(ctx context.Context,
|
||||
taskID, modelName string,
|
||||
) {
|
||||
logger.Infof(ctx, "Starting async download for model: %s, task: %s", modelName, taskID)
|
||||
logger.Infof(ctx, "Starting async download for model, task: %s", taskID)
|
||||
|
||||
// 更新任务状态为下载中
|
||||
h.updateTaskStatus(taskID, "downloading", 0.0, "开始下载模型")
|
||||
@@ -966,16 +964,13 @@ func (h *InitializationHandler) downloadModelAsync(ctx context.Context,
|
||||
h.updateTaskStatus(taskID, "downloading", progress, message)
|
||||
})
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"model_name": modelName,
|
||||
"task_id": taskID,
|
||||
})
|
||||
logger.Error(ctx, "Failed to download model", err)
|
||||
h.updateTaskStatus(taskID, "failed", 0.0, fmt.Sprintf("下载失败: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
// 下载成功
|
||||
logger.Infof(ctx, "Model %s downloaded successfully, task: %s", modelName, taskID)
|
||||
logger.Infof(ctx, "Model downloaded successfully, task: %s", taskID)
|
||||
h.updateTaskStatus(taskID, "completed", 100.0, "下载完成")
|
||||
}
|
||||
|
||||
@@ -993,9 +988,7 @@ func (h *InitializationHandler) pullModelWithProgress(ctx context.Context,
|
||||
// 检查模型是否已存在
|
||||
available, err := h.ollamaService.IsModelAvailable(ctx, modelName)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"model_name": modelName,
|
||||
})
|
||||
logger.Error(ctx, "Failed to check model availability", err)
|
||||
return err
|
||||
}
|
||||
if available {
|
||||
@@ -1003,8 +996,6 @@ func (h *InitializationHandler) pullModelWithProgress(ctx context.Context,
|
||||
return nil
|
||||
}
|
||||
|
||||
logger.GetLogger(ctx).Infof("Pulling model %s...", modelName)
|
||||
|
||||
// 创建下载请求
|
||||
pullReq := &api.PullRequest{
|
||||
Name: modelName,
|
||||
@@ -1026,8 +1017,7 @@ func (h *InitializationHandler) pullModelWithProgress(ctx context.Context,
|
||||
progressCallback(progressPercent, message)
|
||||
|
||||
logger.Infof(ctx,
|
||||
"Download progress for %s: %.2f%% - %s",
|
||||
modelName, progressPercent, message,
|
||||
"Download progress: %.2f%% - %s", progressPercent, message,
|
||||
)
|
||||
return nil
|
||||
})
|
||||
@@ -1062,18 +1052,18 @@ func (h *InitializationHandler) GetCurrentConfigByKB(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
kbIdStr := c.Param("kbId")
|
||||
|
||||
logger.Info(ctx, "Getting configuration for knowledge base", "kbId", kbIdStr)
|
||||
logger.Info(ctx, "Getting configuration for knowledge base")
|
||||
|
||||
// 获取指定知识库信息
|
||||
kb, err := h.kbService.GetKnowledgeBaseByID(ctx, kbIdStr)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{"kbId": kbIdStr})
|
||||
logger.Error(ctx, "Failed to get knowledge base", err)
|
||||
c.Error(errors.NewInternalServerError("获取知识库信息失败: " + err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if kb == nil {
|
||||
logger.Error(ctx, "Knowledge base not found", "kbId", kbIdStr)
|
||||
logger.Error(ctx, "Knowledge base not found")
|
||||
c.Error(errors.NewNotFoundError("知识库不存在"))
|
||||
return
|
||||
}
|
||||
@@ -1090,7 +1080,7 @@ func (h *InitializationHandler) GetCurrentConfigByKB(c *gin.Context) {
|
||||
if modelID != "" {
|
||||
model, err := h.modelService.GetModelByID(ctx, modelID)
|
||||
if err != nil {
|
||||
logger.Warn(ctx, "Failed to get model", "kbId", kbIdStr, "modelId", modelID, "error", err)
|
||||
logger.Warn(ctx, "Failed to get model", err)
|
||||
// 如果模型不存在或获取失败,继续处理其他模型
|
||||
continue
|
||||
}
|
||||
@@ -1114,7 +1104,7 @@ func (h *InitializationHandler) GetCurrentConfigByKB(c *gin.Context) {
|
||||
// 构建配置响应
|
||||
config := h.buildConfigResponse(ctx, models, kb, hasFiles)
|
||||
|
||||
logger.Info(ctx, "Knowledge base configuration retrieved successfully", "kbId", kbIdStr)
|
||||
logger.Info(ctx, "Knowledge base configuration retrieved successfully")
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": config,
|
||||
@@ -1298,12 +1288,7 @@ func (h *InitializationHandler) CheckRemoteModel(c *gin.Context) {
|
||||
// 检查远程模型连接
|
||||
available, message := h.checkRemoteModelConnection(ctx, modelConfig)
|
||||
|
||||
logger.Info(ctx,
|
||||
fmt.Sprintf(
|
||||
"Remote model check completed: modelName=%s, baseUrl=%s, available=%v, message=%s",
|
||||
req.ModelName, req.BaseURL, available, message,
|
||||
),
|
||||
)
|
||||
logger.Infof(ctx, "Remote model check completed, available: %v, message: %s", available, message)
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
@@ -1359,7 +1344,7 @@ func (h *InitializationHandler) TestEmbeddingModel(c *gin.Context) {
|
||||
sample := "hello"
|
||||
vec, err := emb.Embed(ctx, sample)
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{"model": req.ModelName})
|
||||
logger.Error(ctx, "Failed to create embedder", err)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": gin.H{`available`: false, `message`: fmt.Sprintf("调用Embedding失败: %v", err), `dimension`: 0},
|
||||
@@ -1367,7 +1352,7 @@ func (h *InitializationHandler) TestEmbeddingModel(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "Embedding test succeeded, dim=%d", len(vec))
|
||||
logger.Infof(ctx, "Embedding test succeeded, dimension: %d", len(vec))
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": gin.H{`available`: true, `message`: fmt.Sprintf("测试成功,向量维度=%d", len(vec)), `dimension`: len(vec)},
|
||||
@@ -1497,11 +1482,7 @@ func (h *InitializationHandler) CheckRerankModel(c *gin.Context) {
|
||||
ctx, req.ModelName, req.BaseURL, req.APIKey,
|
||||
)
|
||||
|
||||
logger.Info(ctx,
|
||||
fmt.Sprintf("Rerank model check completed: modelName=%s, baseUrl=%s, available=%v, message=%s",
|
||||
req.ModelName, req.BaseURL, available, message,
|
||||
),
|
||||
)
|
||||
logger.Infof(ctx, "Rerank model check completed, available: %v, message: %s", available, message)
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
@@ -1565,8 +1546,6 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
}
|
||||
switch req.StorageType {
|
||||
case "cos":
|
||||
logger.Infof(ctx, "COS config: Region=%s, Bucket=%s, App=%s, Prefix=%s",
|
||||
req.COSRegion, req.COSBucketName, req.COSAppID, req.COSPathPrefix)
|
||||
// 必填:SecretID/SecretKey/Region/BucketName/AppID;PathPrefix 可选
|
||||
if req.COSSecretID == "" || req.COSSecretKey == "" ||
|
||||
req.COSRegion == "" || req.COSBucketName == "" ||
|
||||
@@ -1576,7 +1555,6 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
case "minio":
|
||||
logger.Infof(ctx, "MinIO config: Bucket=%s, PathPrefix=%s", req.MinioBucketName, req.MinioPathPrefix)
|
||||
if req.MinioBucketName == "" {
|
||||
logger.Error(ctx, "MinIO configuration is required")
|
||||
c.Error(errors.NewBadRequestError("MinIO配置信息不能为空"))
|
||||
@@ -1588,9 +1566,6 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "VLM config: Model=%s, URL=%s, HasKey=%v, Type=%s",
|
||||
req.VLMModel, req.VLMBaseURL, req.VLMAPIKey != "", req.VLMInterfaceType)
|
||||
|
||||
// 获取上传的图片文件
|
||||
file, header, err := c.Request.FormFile("image")
|
||||
if err != nil {
|
||||
@@ -1616,12 +1591,14 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
logger.Infof(ctx, "Processing image: %s, size: %d bytes", header.Filename, header.Size)
|
||||
|
||||
// 解析文档分割配置
|
||||
chunkSize, err := strconv.Atoi(req.ChunkSize)
|
||||
chunkSizeInt64, err := strconv.ParseInt(req.ChunkSize, 10, 0)
|
||||
chunkSize := int(chunkSizeInt64)
|
||||
if err != nil || chunkSize < 100 || chunkSize > 10000 {
|
||||
chunkSize = 1000
|
||||
}
|
||||
|
||||
chunkOverlap, err := strconv.Atoi(req.ChunkOverlap)
|
||||
chunkOverlapInt64, err := strconv.ParseInt(req.ChunkOverlap, 10, 0)
|
||||
chunkOverlap := int(chunkOverlapInt64)
|
||||
if err != nil || chunkOverlap < 0 || chunkOverlap >= chunkSize {
|
||||
chunkOverlap = 200
|
||||
}
|
||||
@@ -1653,11 +1630,7 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
processingTime := time.Since(startTime).Milliseconds()
|
||||
|
||||
if err != nil {
|
||||
logger.ErrorWithFields(ctx, err, map[string]interface{}{
|
||||
"vlm_model": req.VLMModel,
|
||||
"vlm_base_url": req.VLMBaseURL,
|
||||
"filename": header.Filename,
|
||||
})
|
||||
logger.Error(ctx, "Failed to test multimodal", err)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": gin.H{
|
||||
@@ -1669,7 +1642,7 @@ func (h *InitializationHandler) TestMultimodalFunction(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Info(ctx, fmt.Sprintf("Multimodal test completed successfully in %dms", processingTime))
|
||||
logger.Infof(ctx, "Multimodal test completed successfully in %dms", processingTime)
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
|
||||
@@ -48,7 +48,6 @@ func (h *MCPServiceHandler) CreateMCPService(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof(ctx, "MCP service created successfully: %s", service.Name)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"success": true,
|
||||
"data": service,
|
||||
|
||||
@@ -110,6 +110,7 @@ func TestRemoteAPIChat(t *testing.T) {
|
||||
t.Run("Basic Chat", func(t *testing.T) {
|
||||
response, err := chat.Chat(ctx, testMessages, testOptions)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, response, "response should not be nil")
|
||||
assert.NotEmpty(t, response.Content)
|
||||
assert.Greater(t, response.Usage.TotalTokens, 0)
|
||||
assert.Greater(t, response.Usage.PromptTokens, 0)
|
||||
|
||||
@@ -147,8 +147,6 @@ func (s *OllamaService) PullModel(ctx context.Context, modelName string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
logger.GetLogger(ctx).Infof("Pulling model %s...", modelName)
|
||||
|
||||
// Use official client to pull model
|
||||
pullReq := &api.PullRequest{
|
||||
Name: modelName,
|
||||
|
||||
Reference in New Issue
Block a user