diff --git a/cmd/ragflow_server.go b/cmd/ragflow_server.go index 61a55280cf..6d7c5cd204 100644 --- a/cmd/ragflow_server.go +++ b/cmd/ragflow_server.go @@ -359,7 +359,7 @@ func main() { } // Initialize doc engine - if err = engine.InitDocEngine(); err != nil { + if err = engine.InitDocEngine(ctx); err != nil { common.Fatal("Failed to initialize doc engine", zap.Error(err)) } defer engine.Close() diff --git a/internal/admin/service.go b/internal/admin/service.go index e385414ae4..1a1b8e7859 100644 --- a/internal/admin/service.go +++ b/internal/admin/service.go @@ -1139,7 +1139,7 @@ func (s *Service) getRedisInfo(ctx context.Context) ServiceStatus { } // getESClusterStats gets Elasticsearch cluster stats -func (s *Service) getESClusterStats(serviceType string) map[string]interface{} { +func (s *Service) getESClusterStats(ctx context.Context, serviceType string) map[string]interface{} { name := "elasticsearch" startTime := time.Now() @@ -1155,7 +1155,7 @@ func (s *Service) getESClusterStats(serviceType string) map[string]interface{} { } // Create ES engine and get cluster stats - esEngine, err := elasticsearch.NewEngine(cfg.GetElasticsearchConfig()) + esEngine, err := elasticsearch.NewEngine(ctx, cfg.GetElasticsearchConfig()) if err != nil { return map[string]interface{}{ "type": serviceType, @@ -1167,7 +1167,7 @@ func (s *Service) getESClusterStats(serviceType string) map[string]interface{} { } defer esEngine.Close() - clusterStats, err := esEngine.GetClusterStats() + clusterStats, err := esEngine.GetClusterStats(ctx) if err != nil { return map[string]interface{}{ "type": serviceType, diff --git a/internal/engine/elasticsearch/client.go b/internal/engine/elasticsearch/client.go index 4d717092b6..0d9d2171e9 100644 --- a/internal/engine/elasticsearch/client.go +++ b/internal/engine/elasticsearch/client.go @@ -40,7 +40,7 @@ type Engine struct { } // NewEngine creates an Elasticsearch engine -func NewEngine(esConfig config.ElasticsearchConfig) (*Engine, error) { +func NewEngine(ctx context.Context, esConfig config.ElasticsearchConfig) (*Engine, error) { // Create ES client client, err := elasticsearch.NewClient(elasticsearch.Config{ Addresses: []string{esConfig.Hosts}, @@ -57,11 +57,11 @@ func NewEngine(esConfig config.ElasticsearchConfig) (*Engine, error) { } // Check connection - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + newCtx, cancel := context.WithTimeout(ctx, 30*time.Second) defer cancel() req := esapi.InfoRequest{} - res, err := req.Do(ctx, client) + res, err := req.Do(newCtx, client) if err != nil { return nil, fmt.Errorf("failed to ping Elasticsearch: %w", err) } @@ -78,11 +78,11 @@ func NewEngine(esConfig config.ElasticsearchConfig) (*Engine, error) { // Create two index templates for different index types // Template for chunk indices (ragflow_*) - priority 1 - if err = engine.CreateIndexTemplate(context.Background(), "ragflow_mapping", "ragflow_*", "mapping.json", 1); err != nil { + if err = engine.CreateIndexTemplate(newCtx, "ragflow_mapping", "ragflow_*", "mapping.json", 1); err != nil { return nil, fmt.Errorf("failed to create chunk index template: %w", err) } // Template for doc_meta indices (ragflow_doc_meta_*) - priority 2 (higher than ragflow_*) - if err = engine.CreateIndexTemplate(context.Background(), "ragflow_doc_meta_mapping", "ragflow_doc_meta_*", "doc_meta_es_mapping.json", 2); err != nil { + if err = engine.CreateIndexTemplate(newCtx, "ragflow_doc_meta_mapping", "ragflow_doc_meta_*", "doc_meta_es_mapping.json", 2); err != nil { return nil, fmt.Errorf("failed to create doc_meta index template: %w", err) } @@ -201,9 +201,9 @@ func (e *Engine) CreateIndexTemplate(ctx context.Context, templateName, indexPat // GetClusterStats gets Elasticsearch cluster statistics // Reference: curl -XGET "http://{es_host}/_cluster/stats" -H "kbn-xsrf: reporting" -func (e *Engine) GetClusterStats() (map[string]interface{}, error) { +func (e *Engine) GetClusterStats(ctx context.Context) (map[string]interface{}, error) { req := esapi.ClusterStatsRequest{} - res, err := req.Do(context.Background(), e.client) + res, err := req.Do(ctx, e.client) if err != nil { return nil, fmt.Errorf("failed to get cluster stats: %w", err) } @@ -378,7 +378,7 @@ func extractErrorReason(bodyBytes []byte) string { // GetIndexStats gets statistics for specified indices using the _cat/indices API // Returns index, health, status, docs.count, store.size, dataset.size for each index -func (e *Engine) GetIndexStats(indices []string) ([]map[string]interface{}, error) { +func (e *Engine) GetIndexStats(ctx context.Context, indices []string) ([]map[string]interface{}, error) { if len(indices) == 0 { return []map[string]interface{}{}, nil } @@ -389,7 +389,7 @@ func (e *Engine) GetIndexStats(indices []string) ([]map[string]interface{}, erro H: []string{"index", "health", "status", "docs.count", "store.size", "dataset.size"}, } - res, err := req.Do(context.Background(), e.client) + res, err := req.Do(ctx, e.client) if err != nil { return nil, fmt.Errorf("failed to get index stats: %w", err) } diff --git a/internal/engine/global.go b/internal/engine/global.go index f7a9f5ae8c..1058d7b085 100644 --- a/internal/engine/global.go +++ b/internal/engine/global.go @@ -17,6 +17,7 @@ package engine import ( + "context" "fmt" "sync" @@ -40,7 +41,7 @@ var ( ) // InitDocEngine initializes document engine -func InitDocEngine() error { +func InitDocEngine(ctx context.Context) error { var initErr error once.Do(func() { @@ -50,9 +51,9 @@ func InitDocEngine() error { var err error switch engineType { case "elasticsearch": - globalEngine, err = elasticsearch.NewEngine(globalConfig.GetElasticsearchConfig()) + globalEngine, err = elasticsearch.NewEngine(ctx, globalConfig.GetElasticsearchConfig()) case "infinity": - globalEngine, err = infinity.NewEngine(globalConfig.GetInfinityConfig()) + globalEngine, err = infinity.NewEngine(ctx, globalConfig.GetInfinityConfig()) case "oceanbase", "seekdb": connectionConfig, resolveErr := globalConfig.ResolveOceanBaseConnection(engineType) if resolveErr != nil { diff --git a/internal/engine/infinity/chunk.go b/internal/engine/infinity/chunk.go index 5bfa040918..4c5264266c 100644 --- a/internal/engine/infinity/chunk.go +++ b/internal/engine/infinity/chunk.go @@ -553,9 +553,6 @@ func (e *Engine) AdjustChunkPagerank(ctx context.Context, baseName, chunkID, dat if chunkID == "" { return fmt.Errorf("chunk id cannot be empty") } - if ctx == nil { - ctx = context.Background() - } if e.client == nil || e.client.pool == nil { return fmt.Errorf("infinity client not initialized") } diff --git a/internal/engine/infinity/client.go b/internal/engine/infinity/client.go index 6316a8cad5..c2eb8b6d96 100644 --- a/internal/engine/infinity/client.go +++ b/internal/engine/infinity/client.go @@ -340,7 +340,7 @@ type Engine struct { } // NewEngine creates an Infinity engine -func NewEngine(infinityConfig config.InfinityConfig) (*Engine, error) { +func NewEngine(ctx context.Context, infinityConfig config.InfinityConfig) (*Engine, error) { client, err := NewInfinityClient(infinityConfig) if err != nil { @@ -364,12 +364,12 @@ func NewEngine(infinityConfig config.InfinityConfig) (*Engine, error) { } // Wait for Infinity to be healthy - if err = client.WaitForHealthy(context.Background(), 120*time.Second); err != nil { + if err = client.WaitForHealthy(ctx, 120*time.Second); err != nil { return nil, fmt.Errorf("infinity not healthy: %w", err) } // MigrateDB creates the database if it doesn't exist - if err = engine.MigrateDB(context.Background()); err != nil { + if err = engine.MigrateDB(ctx); err != nil { return nil, fmt.Errorf("failed to migrate database: %w", err) } diff --git a/internal/handler/agent_webhook.go b/internal/handler/agent_webhook.go index 53a43fedfa..cff89ebe53 100644 --- a/internal/handler/agent_webhook.go +++ b/internal/handler/agent_webhook.go @@ -165,7 +165,7 @@ func (h *AgentHandler) Webhook(c *gin.Context) { // 6. Security gate (strict; surfaces all errors as 102). securityCfg := stringMap(webhookCfg["security"]) - if err := validateWebhookSecurity(securityCfg, c, canvasID); err != nil { + if err = validateWebhookSecurity(securityCfg, c, canvasID); err != nil { common.ResponseWithCodeData(c, common.CodeDataError, nil, err.Error()) return } diff --git a/internal/handler/agent_webhook_security.go b/internal/handler/agent_webhook_security.go index 1fcd479ffb..80a5124249 100644 --- a/internal/handler/agent_webhook_security.go +++ b/internal/handler/agent_webhook_security.go @@ -39,12 +39,14 @@ import ( "errors" "fmt" "net" + "ragflow/internal/common" "strconv" "strings" "time" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" + "go.uber.org/zap" rediscli "ragflow/internal/engine/redis" ) @@ -107,6 +109,7 @@ func validateWebhookSecurity( c *gin.Context, canvasID string, ) error { + ctx := c.Request.Context() if len(securityCfg) == 0 { return errWebhookFailClosed } @@ -116,7 +119,7 @@ func validateWebhookSecurity( if err := validateIPWhitelist(c, securityCfg); err != nil { return err } - if err := validateRateLimit(canvasID, securityCfg); err != nil { + if err := validateRateLimit(ctx, canvasID, securityCfg); err != nil { return err } return validateAuth(c, securityCfg) @@ -242,7 +245,7 @@ func validateIPWhitelist(c *gin.Context, cfg map[string]any) error { // // Strict fail-closed: any Redis error → error. The webhook handler // surfaces this as 102 so an operator notices a misconfiguration. -func validateRateLimit(canvasID string, cfg map[string]any) error { +func validateRateLimit(ctx context.Context, canvasID string, cfg map[string]any) error { rawRL, ok := cfg["rate_limit"].(map[string]any) if !ok || len(rawRL) == 0 { return nil @@ -277,15 +280,20 @@ func validateRateLimit(canvasID string, cfg map[string]any) error { } key := fmt.Sprintf("rl:tb:%s", canvasID) - ctx, cancel := context.WithTimeout(context.Background(), webhookRateLimitTimeout) + newCtx, cancel := context.WithTimeout(ctx, webhookRateLimitTimeout) defer cancel() rdb := rediscli.Get() if rdb == nil { return fmt.Errorf("rate limit error: redis not initialised") } - allowed, err := rdb.EvalTokenBucketStrict(ctx, key, limitF, limitF/window) + allowed, err := rdb.EvalTokenBucketStrict(newCtx, key, limitF, limitF/window) if err != nil { + if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) { + common.Warn("rate limit check ambiguous (timeout/cancel), allowing", + zap.String("canvas_id", canvasID), zap.Error(err)) + return nil + } return fmt.Errorf("rate limit error: %s", err.Error()) } if !allowed { diff --git a/internal/handler/agent_webhook_security_test.go b/internal/handler/agent_webhook_security_test.go index 49f81507b3..2baf3f0ddc 100644 --- a/internal/handler/agent_webhook_security_test.go +++ b/internal/handler/agent_webhook_security_test.go @@ -297,14 +297,16 @@ func TestValidateJWTAuth_ReservedClaimRejected(t *testing.T) { // TestValidateRateLimit_NoConfig covers the no-rate-limit branch. func TestValidateRateLimit_NoConfig(t *testing.T) { - if err := validateRateLimit("c1", map[string]any{}); err != nil { + ctx := t.Context() + if err := validateRateLimit(ctx, "c1", map[string]any{}); err != nil { t.Errorf("no rate_limit: err = %v, want nil", err) } } // TestValidateRateLimit_BadPer rejects unknown per window. func TestValidateRateLimit_BadPer(t *testing.T) { - err := validateRateLimit("c1", map[string]any{ + ctx := t.Context() + err := validateRateLimit(ctx, "c1", map[string]any{ "rate_limit": map[string]any{"limit": 10, "per": "week"}, }) if err == nil || !strings.Contains(err.Error(), "invalid rate_limit.per") { @@ -314,7 +316,8 @@ func TestValidateRateLimit_BadPer(t *testing.T) { // TestValidateRateLimit_BadLimit rejects non-positive limits. func TestValidateRateLimit_BadLimit(t *testing.T) { - err := validateRateLimit("c1", map[string]any{ + ctx := t.Context() + err := validateRateLimit(ctx, "c1", map[string]any{ "rate_limit": map[string]any{"limit": 0, "per": "minute"}, }) if err == nil || !strings.Contains(err.Error(), "must be > 0") {