From 2f45bf7d86e4c994ec2dc0704b7a07fb187c09d2 Mon Sep 17 00:00:00 2001 From: wizardchen Date: Tue, 25 Nov 2025 22:06:43 +0800 Subject: [PATCH] refactor: Improve Elasticsearch query construction and error handling --- .../retriever/elasticsearch/v7/repository.go | 117 +++++++++++++----- internal/application/service/graph.go | 7 ++ internal/application/service/knowledgebase.go | 1 - internal/application/service/mcp_service.go | 1 - internal/application/service/session.go | 53 ++++---- internal/application/service/user.go | 17 ++- internal/handler/chunk.go | 26 +--- internal/handler/evaluation.go | 9 +- internal/handler/initialization.go | 77 ++++-------- internal/handler/mcp_service.go | 1 - internal/models/chat/remote_api_test.go | 1 + internal/models/utils/ollama/ollama.go | 2 - 12 files changed, 158 insertions(+), 154 deletions(-) diff --git a/internal/application/repository/retriever/elasticsearch/v7/repository.go b/internal/application/repository/retriever/elasticsearch/v7/repository.go index 7c9387301..0e282ceb2 100644 --- a/internal/application/repository/retriever/elasticsearch/v7/repository.go +++ b/internal/application/repository/retriever/elasticsearch/v7/repository.go @@ -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 } diff --git a/internal/application/service/graph.go b/internal/application/service/graph.go index c627c5565..644275d1c 100644 --- a/internal/application/service/graph.go +++ b/internal/application/service/graph.go @@ -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) { diff --git a/internal/application/service/knowledgebase.go b/internal/application/service/knowledgebase.go index c8cb1cbfb..9a673c14f 100644 --- a/internal/application/service/knowledgebase.go +++ b/internal/application/service/knowledgebase.go @@ -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) diff --git a/internal/application/service/mcp_service.go b/internal/application/service/mcp_service.go index 6891859ad..5cbf9768f 100644 --- a/internal/application/service/mcp_service.go +++ b/internal/application/service/mcp_service.go @@ -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 } diff --git a/internal/application/service/session.go b/internal/application/service/session.go index d84d6f737..537c57af9 100644 --- a/internal/application/service/session.go +++ b/internal/application/service/session.go @@ -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 diff --git a/internal/application/service/user.go b/internal/application/service/user.go index 3ed06a61c..49ef3d61a 100644 --- a/internal/application/service/user.go +++ b/internal/application/service/user.go @@ -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 { diff --git a/internal/handler/chunk.go b/internal/handler/chunk.go index 1f7e52356..eda045ea7 100644 --- a/internal/handler/chunk.go +++ b/internal/handler/chunk.go @@ -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", diff --git a/internal/handler/evaluation.go b/internal/handler/evaluation.go index ad25dc37f..324fda765 100644 --- a/internal/handler/evaluation.go +++ b/internal/handler/evaluation.go @@ -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) diff --git a/internal/handler/initialization.go b/internal/handler/initialization.go index 8666eaac4..9b0f6ebf5 100644 --- a/internal/handler/initialization.go +++ b/internal/handler/initialization.go @@ -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, diff --git a/internal/handler/mcp_service.go b/internal/handler/mcp_service.go index 65164f5e1..df3e5f6d9 100644 --- a/internal/handler/mcp_service.go +++ b/internal/handler/mcp_service.go @@ -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, diff --git a/internal/models/chat/remote_api_test.go b/internal/models/chat/remote_api_test.go index 115ad23e3..782d41f84 100644 --- a/internal/models/chat/remote_api_test.go +++ b/internal/models/chat/remote_api_test.go @@ -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) diff --git a/internal/models/utils/ollama/ollama.go b/internal/models/utils/ollama/ollama.go index df9634c21..b166de5db 100644 --- a/internal/models/utils/ollama/ollama.go +++ b/internal/models/utils/ollama/ollama.go @@ -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,