diff --git a/services/llm-api/internal/domain/model/model_catalog_service.go b/services/llm-api/internal/domain/model/model_catalog_service.go index 08ab3a11..6652840f 100644 --- a/services/llm-api/internal/domain/model/model_catalog_service.go +++ b/services/llm-api/internal/domain/model/model_catalog_service.go @@ -284,8 +284,10 @@ func buildModelCatalogFromModel(provider *Provider, model chat.Model) *ModelCata "temperature", "max_tokens", "top_p", + "top_k", "frequency_penalty", "presence_penalty", + "repetition_penalty", "stop", "stream", "n", diff --git a/services/llm-api/internal/domain/model/provider_model_service.go b/services/llm-api/internal/domain/model/provider_model_service.go index 65673ee5..8dd84c86 100644 --- a/services/llm-api/internal/domain/model/provider_model_service.go +++ b/services/llm-api/internal/domain/model/provider_model_service.go @@ -209,6 +209,19 @@ func (s *ProviderModelService) BatchUpdateActive(ctx context.Context, filter Pro return rowsAffected, nil } +// FindCatalogByID returns a model catalog by ID (used for applying defaults). +func (s *ProviderModelService) FindCatalogByID(ctx context.Context, id uint) (*ModelCatalog, error) { + if id == 0 { + return nil, platformerrors.NewError(ctx, platformerrors.LayerDomain, platformerrors.ErrorTypeValidation, "model catalog ID is required", nil, "c9bde6a4-6bd1-4f8e-97df-e7b03c3f6f73") + } + + catalog, err := s.modelCatalogRepo.FindByID(ctx, id) + if err != nil { + return nil, platformerrors.AsError(ctx, platformerrors.LayerDomain, err, "failed to find model catalog by ID") + } + return catalog, nil +} + func (s *ProviderModelService) BatchUpdateModelDisplayName(ctx context.Context, filter ProviderModelFilter, modelDisplayName string) (int64, error) { rowsAffected, err := s.providerModelRepo.BatchUpdateModelDisplayName(ctx, filter, modelDisplayName) if err != nil { diff --git a/services/llm-api/internal/interfaces/httpserver/handlers/chathandler/chat_handler.go b/services/llm-api/internal/interfaces/httpserver/handlers/chathandler/chat_handler.go index 70a3e613..05d2cda5 100644 --- a/services/llm-api/internal/interfaces/httpserver/handlers/chathandler/chat_handler.go +++ b/services/llm-api/internal/interfaces/httpserver/handlers/chathandler/chat_handler.go @@ -13,6 +13,7 @@ import ( "go.opentelemetry.io/otel/codes" "jan-server/services/llm-api/internal/domain/conversation" + domainmodel "jan-server/services/llm-api/internal/domain/model" "jan-server/services/llm-api/internal/domain/project" "jan-server/services/llm-api/internal/domain/prompt" "jan-server/services/llm-api/internal/domain/usersettings" @@ -28,6 +29,8 @@ import ( "jan-server/services/llm-api/internal/utils/httpclients/chat" "jan-server/services/llm-api/internal/utils/idgen" "jan-server/services/llm-api/internal/utils/platformerrors" + + "github.com/shopspring/decimal" ) const ConversationReferrerContextKey = "conversation_referrer" @@ -201,6 +204,16 @@ func (h *ChatHandler) CreateChatCompletion( // Override the request model with the provider's original model ID request.Model = selectedProviderModel.ProviderOriginalModelID + // Optionally load model catalog (used later to apply default parameters) + var modelCatalog *domainmodel.ModelCatalog + if selectedProviderModel.ModelCatalogID != nil { + modelCatalog, err = h.providerHandler.GetModelCatalogByID(ctx, *selectedProviderModel.ModelCatalogID) + if err != nil { + log := logger.GetLogger() + log.Warn().Err(err).Uint("model_catalog_id", *selectedProviderModel.ModelCatalogID).Msg("failed to load model catalog defaults") + } + } + // Resolve jan_* media placeholders (best-effort) request.Messages = h.resolveMediaPlaceholders(ctx, reqCtx, request.Messages) @@ -283,12 +296,20 @@ func (h *ChatHandler) CreateChatCompletion( } // Handle streaming vs non-streaming + llmRequest := chat.CompletionRequest{ + ChatCompletionRequest: request.ChatCompletionRequest, + TopK: request.TopK, + RepetitionPenalty: request.RepetitionPenalty, + } + if modelCatalog != nil { + h.applyModelDefaultsFromCatalog(&llmRequest, modelCatalog) + } observability.AddSpanEvent(ctx, "calling_llm") llmStartTime := time.Now() if request.Stream { - response, err = h.streamCompletion(ctx, reqCtx, chatClient, conv, request.ChatCompletionRequest) + response, err = h.streamCompletion(ctx, reqCtx, chatClient, conv, llmRequest) } else { - response, err = h.callCompletion(ctx, chatClient, request.ChatCompletionRequest) + response, err = h.callCompletion(ctx, chatClient, llmRequest) } llmDuration := time.Since(llmStartTime) @@ -419,7 +440,7 @@ func (h *ChatHandler) CreateChatCompletion( func (h *ChatHandler) callCompletion( ctx context.Context, chatClient *chat.ChatCompletionClient, - request openai.ChatCompletionRequest, + request chat.CompletionRequest, ) (*openai.ChatCompletionResponse, error) { chatCompletion, err := chatClient.CreateChatCompletion(ctx, "", request) if err != nil { @@ -435,7 +456,7 @@ func (h *ChatHandler) streamCompletion( reqCtx *gin.Context, chatClient *chat.ChatCompletionClient, conv *conversation.Conversation, - request openai.ChatCompletionRequest, + request chat.CompletionRequest, ) (*openai.ChatCompletionResponse, error) { // Create callback to send conversation data before [DONE] var beforeDoneCallback chat.BeforeDoneCallback @@ -525,6 +546,69 @@ func (h *ChatHandler) resolveMediaPlaceholders(ctx context.Context, reqCtx *gin. return messages } +// applyModelDefaultsFromCatalog fills in missing request parameters using defaults from the model catalog. +func (h *ChatHandler) applyModelDefaultsFromCatalog(req *chat.CompletionRequest, catalog *domainmodel.ModelCatalog) { + if req == nil || catalog == nil { + return + } + + defaults := catalog.SupportedParameters.Default + if len(defaults) == 0 { + return + } + + if req.Temperature == 0 { + if val, ok := decimalToFloat32(defaults["temperature"]); ok { + req.Temperature = val + } + } + if req.TopP == 0 { + if val, ok := decimalToFloat32(defaults["top_p"]); ok { + req.TopP = val + } + } + if req.PresencePenalty == 0 { + if val, ok := decimalToFloat32(defaults["presence_penalty"]); ok { + req.PresencePenalty = val + } + } + if req.FrequencyPenalty == 0 { + if val, ok := decimalToFloat32(defaults["frequency_penalty"]); ok { + req.FrequencyPenalty = val + } + } + if req.MaxTokens == 0 { + if val, ok := decimalToInt(defaults["max_tokens"]); ok { + req.MaxTokens = val + } + } + if req.TopK == nil || (req.TopK != nil && *req.TopK == 0) { + if val, ok := decimalToInt(defaults["top_k"]); ok { + req.TopK = &val + } + } + if req.RepetitionPenalty == nil || (req.RepetitionPenalty != nil && *req.RepetitionPenalty == 0) { + if val, ok := decimalToFloat32(defaults["repetition_penalty"]); ok { + req.RepetitionPenalty = &val + } + } +} + +func decimalToFloat32(val *decimal.Decimal) (float32, bool) { + if val == nil { + return 0, false + } + f, _ := val.Float64() + return float32(f), true +} + +func decimalToInt(val *decimal.Decimal) (int, bool) { + if val == nil { + return 0, false + } + return int(val.IntPart()), true +} + // getProjectInstruction loads the project instruction for the conversation, falling back to the stored snapshot. func (h *ChatHandler) getProjectInstruction(ctx context.Context, userID uint, conv *conversation.Conversation) string { if conv == nil || h.projectService == nil { diff --git a/services/llm-api/internal/interfaces/httpserver/handlers/modelhandler/provider_handler.go b/services/llm-api/internal/interfaces/httpserver/handlers/modelhandler/provider_handler.go index f62d503a..a2725c61 100644 --- a/services/llm-api/internal/interfaces/httpserver/handlers/modelhandler/provider_handler.go +++ b/services/llm-api/internal/interfaces/httpserver/handlers/modelhandler/provider_handler.go @@ -257,6 +257,11 @@ func (h *ProviderHandler) UpdateProvider( return &response, nil } +// GetModelCatalogByID returns catalog details (used to apply default parameters). +func (providerHandler *ProviderHandler) GetModelCatalogByID(ctx context.Context, id uint) (*domainmodel.ModelCatalog, error) { + return providerHandler.providerModelService.FindCatalogByID(ctx, id) +} + func (h *ProviderHandler) DeleteProvider(ctx context.Context, publicID string) error { if strings.TrimSpace(publicID) == "" { return platformerrors.NewError(ctx, platformerrors.LayerHandler, platformerrors.ErrorTypeValidation, "provider public ID is required", nil, "0c3f68da-0aa4-4a7c-9cec-c22d47c86f8b") diff --git a/services/llm-api/internal/interfaces/httpserver/requests/chat/chat.go b/services/llm-api/internal/interfaces/httpserver/requests/chat/chat.go index c3a016e7..7d438908 100644 --- a/services/llm-api/internal/interfaces/httpserver/requests/chat/chat.go +++ b/services/llm-api/internal/interfaces/httpserver/requests/chat/chat.go @@ -12,6 +12,9 @@ import ( type ChatCompletionRequest struct { openai.ChatCompletionRequest + TopK *int `json:"top_k,omitempty"` + RepetitionPenalty *float32 `json:"repetition_penalty,omitempty"` + // Conversation can be either a string (conversation ID) or a conversation object // Items from this conversation are prepended to Messages for this response request. // Input items and output items from this response are automatically added to this conversation after completion. diff --git a/services/llm-api/internal/utils/httpclients/chat/chat_completion_client.go b/services/llm-api/internal/utils/httpclients/chat/chat_completion_client.go index 493b6bfc..8da1872e 100644 --- a/services/llm-api/internal/utils/httpclients/chat/chat_completion_client.go +++ b/services/llm-api/internal/utils/httpclients/chat/chat_completion_client.go @@ -78,6 +78,13 @@ type ChatCompletionClient struct { name string } +// CompletionRequest extends the OpenAI chat request with provider-specific fields. +type CompletionRequest struct { + openai.ChatCompletionRequest + TopK *int `json:"top_k,omitempty"` + RepetitionPenalty *float32 `json:"repetition_penalty,omitempty"` +} + type functionCallAccumulator struct { Name string Arguments string @@ -103,7 +110,7 @@ func NewChatCompletionClient(client *resty.Client, name, baseURL string) *ChatCo } } -func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey string, request openai.ChatCompletionRequest) (*openai.ChatCompletionResponse, error) { +func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey string, request CompletionRequest) (*openai.ChatCompletionResponse, error) { // Start OpenTelemetry span for tracking ctx, span := otel.Tracer("chat-completion-client").Start(ctx, "CreateChatCompletion", trace.WithSpanKind(trace.SpanKindClient), @@ -126,9 +133,41 @@ func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey if request.TopP != 0 { span.SetAttributes(attribute.Float64("llm.top_p", float64(request.TopP))) } + if request.PresencePenalty != 0 { + span.SetAttributes(attribute.Float64("llm.presence_penalty", float64(request.PresencePenalty))) + } + if request.FrequencyPenalty != 0 { + span.SetAttributes(attribute.Float64("llm.frequency_penalty", float64(request.FrequencyPenalty))) + } start := time.Now() + // Debug logging: Log request details before sending + log := logger.GetLogger() + topK := 0 + if request.TopK != nil { + topK = *request.TopK + } + repetitionPenalty := float32(0) + if request.RepetitionPenalty != nil { + repetitionPenalty = *request.RepetitionPenalty + } + log.Info(). + Str("provider", c.name). + Str("model", request.Model). + Int("messages", len(request.Messages)). + Bool("stream", request.Stream). + Float32("temperature", request.Temperature). + Int("max_tokens", request.MaxTokens). + Float32("top_p", request.TopP). + Int("top_k", topK). + Float32("presence_penalty", request.PresencePenalty). + Float32("frequency_penalty", request.FrequencyPenalty). + Float32("repetition_penalty", repetitionPenalty). + Msg("[ChatCompletion] Sending request to inference server") + + log.Info().Interface("request", request).Msg("[ChatCompletion] Request body") + var respBody openai.ChatCompletionResponse resp, err := c.prepareRequest(ctx, apiKey). SetBody(request). @@ -181,7 +220,7 @@ func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey return &respBody, nil } -func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (io.ReadCloser, error) { +func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, apiKey string, request CompletionRequest, opts ...StreamOption) (io.ReadCloser, error) { resp, err := c.doStreamingRequest(ctx, apiKey, request, opts...) if err != nil { return nil, err @@ -208,10 +247,10 @@ func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, a } func (c *ChatCompletionClient) StreamChatCompletionToContext(reqCtx *gin.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (*openai.ChatCompletionResponse, error) { - return c.StreamChatCompletionToContextWithCallback(reqCtx, apiKey, request, nil, opts...) + return c.StreamChatCompletionToContextWithCallback(reqCtx, apiKey, CompletionRequest{ChatCompletionRequest: request}, nil, opts...) } -func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx *gin.Context, apiKey string, request openai.ChatCompletionRequest, beforeDone BeforeDoneCallback, opts ...StreamOption) (*openai.ChatCompletionResponse, error) { +func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx *gin.Context, apiKey string, request CompletionRequest, beforeDone BeforeDoneCallback, opts ...StreamOption) (*openai.ChatCompletionResponse, error) { // Start OpenTelemetry span for tracking streaming completion ctx := reqCtx.Request.Context() ctx, span := otel.Tracer("chat-completion-client").Start(ctx, "StreamChatCompletion", @@ -235,9 +274,29 @@ func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx if request.TopP != 0 { span.SetAttributes(attribute.Float64("llm.top_p", float64(request.TopP))) } + if request.PresencePenalty != 0 { + span.SetAttributes(attribute.Float64("llm.presence_penalty", float64(request.PresencePenalty))) + } + if request.FrequencyPenalty != 0 { + span.SetAttributes(attribute.Float64("llm.frequency_penalty", float64(request.FrequencyPenalty))) + } start := time.Now() + // Debug logging: Log request details before sending + log := logger.GetLogger() + log.Debug(). + Str("provider", c.name). + Str("model", request.Model). + Int("messages", len(request.Messages)). + Bool("stream", request.Stream). + Float32("temperature", request.Temperature). + Int("max_tokens", request.MaxTokens). + Float32("top_p", request.TopP). + Float32("presence_penalty", request.PresencePenalty). + Float32("frequency_penalty", request.FrequencyPenalty). + Msg("[StreamChatCompletion] Sending request to inference server") + // force to true to collect tokens request.StreamOptions = &openai.StreamOptions{ IncludeUsage: true, @@ -470,7 +529,33 @@ func (c *ChatCompletionClient) errorFromResponse(ctx context.Context, resp *rest return platformerrors.NewError(ctx, platformerrors.LayerDomain, platformerrors.ErrorTypeExternal, fmt.Sprintf("%s: %s", message, trimmed), nil, "a1f46e0d-4017-4411-ac05-987946c3066d") } -func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (*resty.Response, error) { +func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey string, request CompletionRequest, opts ...StreamOption) (*resty.Response, error) { + // Debug logging: Log request details for streaming calls + log := logger.GetLogger() + topK := 0 + if request.TopK != nil { + topK = *request.TopK + } + repetitionPenalty := float32(0) + if request.RepetitionPenalty != nil { + repetitionPenalty = *request.RepetitionPenalty + } + log.Info(). + Str("provider", c.name). + Str("model", request.Model). + Int("messages", len(request.Messages)). + Bool("stream", request.Stream). + Float32("temperature", request.Temperature). + Int("max_tokens", request.MaxTokens). + Float32("top_p", request.TopP). + Int("top_k", topK). + Float32("presence_penalty", request.PresencePenalty). + Float32("frequency_penalty", request.FrequencyPenalty). + Float32("repetition_penalty", repetitionPenalty). + Msg("[ChatCompletion][Stream] Sending request to inference server") + + log.Info().Interface("request", request).Msg("[ChatCompletion][Stream] Request body") + req := c.prepareRequest(ctx, apiKey). SetBody(request). SetDoNotParseResponse(true) @@ -500,7 +585,7 @@ func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey st return resp, nil } -func (c *ChatCompletionClient) streamResponseToChannel(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, dataChan chan<- string, errChan chan<- error, wg *sync.WaitGroup, opts []StreamOption) { +func (c *ChatCompletionClient) streamResponseToChannel(ctx context.Context, apiKey string, request CompletionRequest, dataChan chan<- string, errChan chan<- error, wg *sync.WaitGroup, opts []StreamOption) { defer wg.Done() resp, err := c.doStreamingRequest(ctx, apiKey, request, opts...) @@ -640,7 +725,7 @@ func (c *ChatCompletionClient) handleStreamingToolCall(toolCall *openai.ToolCall } } -func (c *ChatCompletionClient) buildCompleteResponse(content string, reasoning string, functionCallAccumulator map[int]*functionCallAccumulator, toolCallAccumulator map[int]*toolCallAccumulator, model string, request openai.ChatCompletionRequest) openai.ChatCompletionResponse { +func (c *ChatCompletionClient) buildCompleteResponse(content string, reasoning string, functionCallAccumulator map[int]*functionCallAccumulator, toolCallAccumulator map[int]*toolCallAccumulator, model string, request CompletionRequest) openai.ChatCompletionResponse { message := openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleAssistant, Content: content, diff --git a/services/mcp-tools/internal/infrastructure/search/circuit_breaker.go b/services/mcp-tools/internal/infrastructure/search/circuit_breaker.go new file mode 100644 index 00000000..3aad43e9 --- /dev/null +++ b/services/mcp-tools/internal/infrastructure/search/circuit_breaker.go @@ -0,0 +1,186 @@ +package search + +import ( + "fmt" + "sync" + "time" + + "github.com/rs/zerolog/log" +) + +// CircuitState represents the state of a circuit breaker +type CircuitState int + +const ( + StateClosed CircuitState = iota + StateOpen + StateHalfOpen +) + +func (s CircuitState) String() string { + switch s { + case StateClosed: + return "closed" + case StateOpen: + return "open" + case StateHalfOpen: + return "half-open" + default: + return "unknown" + } +} + +// CircuitBreakerConfig defines circuit breaker behavior +type CircuitBreakerConfig struct { + FailureThreshold int // Number of failures before opening + SuccessThreshold int // Number of successes to close from half-open + Timeout time.Duration // How long to stay open before trying half-open + MaxHalfOpenCalls int // Max concurrent calls in half-open state +} + +// DefaultCircuitBreakerConfig returns sensible defaults +func DefaultCircuitBreakerConfig() CircuitBreakerConfig { + return CircuitBreakerConfig{ + FailureThreshold: 5, + SuccessThreshold: 2, + Timeout: 30 * time.Second, + MaxHalfOpenCalls: 1, + } +} + +// CircuitBreaker implements the circuit breaker pattern +type CircuitBreaker struct { + cfg CircuitBreakerConfig + mu sync.RWMutex + + state CircuitState + failures int + successes int + lastFailureTime time.Time + halfOpenCalls int +} + +// NewCircuitBreaker creates a new circuit breaker +func NewCircuitBreaker(cfg CircuitBreakerConfig) *CircuitBreaker { + return &CircuitBreaker{ + cfg: cfg, + state: StateClosed, + } +} + +// Execute runs a function with circuit breaker protection +func (cb *CircuitBreaker) Execute(operation string, fn func() error) error { + if !cb.allowRequest() { + return fmt.Errorf("circuit breaker is open for %s", operation) + } + + err := fn() + cb.recordResult(operation, err) + return err +} + +// allowRequest determines if a request should be allowed +func (cb *CircuitBreaker) allowRequest() bool { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case StateClosed: + return true + case StateOpen: + // Check if timeout has elapsed + if time.Since(cb.lastFailureTime) > cb.cfg.Timeout { + log.Info().Msg("circuit breaker transitioning to half-open") + cb.state = StateHalfOpen + cb.halfOpenCalls = 0 + return true + } + return false + case StateHalfOpen: + // Limit concurrent calls in half-open state + if cb.halfOpenCalls < cb.cfg.MaxHalfOpenCalls { + cb.halfOpenCalls++ + return true + } + return false + default: + return false + } +} + +// recordResult updates circuit breaker state based on result +func (cb *CircuitBreaker) recordResult(operation string, err error) { + cb.mu.Lock() + defer cb.mu.Unlock() + + if err != nil { + cb.failures++ + cb.successes = 0 + cb.lastFailureTime = time.Now() + + if cb.state == StateHalfOpen { + log.Warn(). + Str("operation", operation). + Msg("circuit breaker opening from half-open due to failure") + cb.state = StateOpen + cb.halfOpenCalls = 0 + } else if cb.state == StateClosed && cb.failures >= cb.cfg.FailureThreshold { + log.Warn(). + Str("operation", operation). + Int("failures", cb.failures). + Msg("circuit breaker opening due to failure threshold") + cb.state = StateOpen + } + } else { + cb.successes++ + + if cb.state == StateHalfOpen { + if cb.successes >= cb.cfg.SuccessThreshold { + log.Info(). + Str("operation", operation). + Int("successes", cb.successes). + Msg("circuit breaker closing from half-open") + cb.state = StateClosed + cb.failures = 0 + cb.successes = 0 + cb.halfOpenCalls = 0 + } + } else if cb.state == StateClosed { + // Reset failure count on success + cb.failures = 0 + } + } +} + +// GetState returns the current circuit breaker state +func (cb *CircuitBreaker) GetState() CircuitState { + cb.mu.RLock() + defer cb.mu.RUnlock() + return cb.state +} + +// GetMetrics returns current circuit breaker metrics +func (cb *CircuitBreaker) GetMetrics() map[string]any { + cb.mu.RLock() + defer cb.mu.RUnlock() + + return map[string]any{ + "state": cb.state.String(), + "failures": cb.failures, + "successes": cb.successes, + "last_failure_time": cb.lastFailureTime, + "half_open_calls": cb.halfOpenCalls, + } +} + +// Reset manually resets the circuit breaker to closed state +func (cb *CircuitBreaker) Reset() { + cb.mu.Lock() + defer cb.mu.Unlock() + + log.Info().Msg("manually resetting circuit breaker") + cb.state = StateClosed + cb.failures = 0 + cb.successes = 0 + cb.halfOpenCalls = 0 +} diff --git a/services/mcp-tools/internal/infrastructure/search/client.go b/services/mcp-tools/internal/infrastructure/search/client.go index 71fddd5d..c8e47ea5 100644 --- a/services/mcp-tools/internal/infrastructure/search/client.go +++ b/services/mcp-tools/internal/infrastructure/search/client.go @@ -46,6 +46,9 @@ type SearchClient struct { serperClient *resty.Client fallbackClient *resty.Client searxClient *resty.Client + retryConfig RetryConfig + serperCB *CircuitBreaker + searxCB *CircuitBreaker } var _ domainsearch.SearchClient = (*SearchClient)(nil) @@ -80,6 +83,9 @@ func NewSearchClient(cfg ClientConfig) *SearchClient { serperClient: serperHTTP, fallbackClient: fallbackHTTP, searxClient: searxHTTP, + retryConfig: DefaultRetryConfig(), + serperCB: NewCircuitBreaker(DefaultCircuitBreakerConfig()), + searxCB: NewCircuitBreaker(DefaultCircuitBreakerConfig()), } } @@ -176,6 +182,12 @@ func (c *SearchClient) resolveOfflineMode(override *bool) bool { } func (c *SearchClient) searchViaSerper(ctx context.Context, query domainsearch.SearchRequest) (*domainsearch.SearchResponse, error) { + // Check circuit breaker + if c.serperCB.GetState() == StateOpen { + log.Warn().Msg("serper circuit breaker is open, skipping") + return nil, fmt.Errorf("serper circuit breaker is open") + } + body := map[string]any{ "q": query.Q, } @@ -203,21 +215,43 @@ func (c *SearchClient) searchViaSerper(ctx context.Context, query domainsearch.S body["tbs"] = string(*query.TBS) } - var result domainsearch.SearchResponse - resp, err := c.serperClient.R(). - SetContext(ctx). - SetHeader("X-API-KEY", c.cfg.SerperAPIKey). - SetHeader("Content-Type", "application/json"). - SetBody(body). - SetResult(&result). - Post(serperSearchEndpoint) + var result *domainsearch.SearchResponse + + // Retry with exponential backoff + resultPtr, err := WithRetry(ctx, c.retryConfig, "serper_search", func() (*domainsearch.SearchResponse, error) { + var res domainsearch.SearchResponse + resp, err := c.serperClient.R(). + SetContext(ctx). + SetHeader("X-API-KEY", c.cfg.SerperAPIKey). + SetHeader("Content-Type", "application/json"). + SetBody(body). + SetResult(&res). + Post(serperSearchEndpoint) + + if err != nil { + return nil, fmt.Errorf("failed to query Serper search API: %w", err) + } + if resp.IsError() { + return nil, fmt.Errorf("Serper search API error (status %d): %s", resp.StatusCode(), resp.String()) + } + + return &res, nil + }) + + // Update circuit breaker + c.serperCB.recordResult("serper_search", err) + if err != nil { - return nil, fmt.Errorf("failed to query Serper search API: %w", err) + return nil, err } - - if resp.IsError() { - return nil, fmt.Errorf("Serper search API error (status %d): %s", resp.StatusCode(), resp.String()) + + result = resultPtr + + // Validate response + if validationErr := ValidateSearchResponse(result, 0); validationErr != nil { + log.Warn().Err(validationErr).Msg("serper search returned invalid response") + return EnrichEmptyResponse(result, query.Q, "validation_failed"), nil } if result.SearchParameters == nil { @@ -230,7 +264,7 @@ func (c *SearchClient) searchViaSerper(ctx context.Context, query domainsearch.S result.SearchParameters["location_hint"] = *query.LocationHint } - return &result, nil + return result, nil } func (c *SearchClient) searchViaSearxng(ctx context.Context, query domainsearch.SearchRequest) (*domainsearch.SearchResponse, error) { @@ -238,35 +272,55 @@ func (c *SearchClient) searchViaSearxng(ctx context.Context, query domainsearch. return nil, fmt.Errorf("searxng client not configured") } - req := c.searxClient.R(). - SetContext(ctx). - SetQueryParam("q", query.Q). - SetQueryParam("format", "json"). - SetQueryParam("safesearch", "1") - - if query.HL != nil { - req.SetQueryParam("language", *query.HL) + // Check circuit breaker + if c.searxCB.GetState() == StateOpen { + log.Warn().Msg("searxng circuit breaker is open, skipping") + return nil, fmt.Errorf("searxng circuit breaker is open") } - if query.Page != nil && *query.Page > 1 { - req.SetQueryParam("p", strconv.Itoa(*query.Page)) - } - if query.Num != nil && *query.Num > 0 { - req.SetQueryParam("num", strconv.Itoa(*query.Num)) - } - if query.TBS != nil { - if mapped := mapTBSToSearxng(*query.TBS); mapped != "" { - req.SetQueryParam("time_range", mapped) + + // Retry with exponential backoff + resultPtr, err := WithRetry(ctx, c.retryConfig, "searxng_search", func() (*searxngResponse, error) { + req := c.searxClient.R(). + SetContext(ctx). + SetQueryParam("q", query.Q). + SetQueryParam("format", "json"). + SetQueryParam("safesearch", "1") + + if query.HL != nil { + req.SetQueryParam("language", *query.HL) + } + if query.Page != nil && *query.Page > 1 { + req.SetQueryParam("p", strconv.Itoa(*query.Page)) + } + if query.Num != nil && *query.Num > 0 { + req.SetQueryParam("num", strconv.Itoa(*query.Num)) + } + if query.TBS != nil { + if mapped := mapTBSToSearxng(*query.TBS); mapped != "" { + req.SetQueryParam("time_range", mapped) + } } - } - var result searxngResponse - resp, err := req.SetResult(&result).Get(searxngSearchPath) + var result searxngResponse + resp, err := req.SetResult(&result).Get(searxngSearchPath) + if err != nil { + return nil, fmt.Errorf("failed to query SearXNG API: %w", err) + } + if resp.IsError() { + return nil, fmt.Errorf("SearXNG API error (status %d): %s", resp.StatusCode(), resp.String()) + } + + return &result, nil + }) + + // Update circuit breaker + c.searxCB.recordResult("searxng_search", err) + if err != nil { - return nil, fmt.Errorf("failed to query SearXNG API: %w", err) - } - if resp.IsError() { - return nil, fmt.Errorf("SearXNG API error (status %d): %s", resp.StatusCode(), resp.String()) + return nil, err } + + result := *resultPtr limit := 10 if query.Num != nil && *query.Num > 0 { @@ -297,10 +351,18 @@ func (c *SearchClient) searchViaSearxng(ctx context.Context, query domainsearch. searchMetadata["location_hint"] = *query.LocationHint } - return &domainsearch.SearchResponse{ + searchResp := &domainsearch.SearchResponse{ SearchParameters: searchMetadata, Organic: results, - }, nil + } + + // Validate response + if validationErr := ValidateSearchResponse(searchResp, 0); validationErr != nil { + log.Warn().Err(validationErr).Msg("searxng search returned invalid response") + return EnrichEmptyResponse(searchResp, query.Q, "validation_failed"), nil + } + + return searchResp, nil } func mapTBSToSearxng(t domainsearch.TBSTimeRange) string { @@ -321,6 +383,12 @@ func mapTBSToSearxng(t domainsearch.TBSTimeRange) string { } func (c *SearchClient) fetchViaSerper(ctx context.Context, query domainsearch.FetchWebpageRequest) (*domainsearch.FetchWebpageResponse, error) { + // Check circuit breaker + if c.serperCB.GetState() == StateOpen { + log.Warn().Msg("serper circuit breaker is open for scraping, using fallback") + return nil, fmt.Errorf("serper circuit breaker is open") + } + body := map[string]any{ "url": query.Url, } @@ -328,54 +396,90 @@ func (c *SearchClient) fetchViaSerper(ctx context.Context, query domainsearch.Fe body["includeMarkdown"] = *query.IncludeMarkdown } - var result domainsearch.FetchWebpageResponse - resp, err := c.serperClient.R(). - SetContext(ctx). - SetHeader("X-API-KEY", c.cfg.SerperAPIKey). - SetHeader("Content-Type", "application/json"). - SetBody(body). - SetResult(&result). - Post(serperScrapeEndpoint) + // Retry with exponential backoff + result, err := WithRetry(ctx, c.retryConfig, "serper_scrape", func() (*domainsearch.FetchWebpageResponse, error) { + var res domainsearch.FetchWebpageResponse + resp, err := c.serperClient.R(). + SetContext(ctx). + SetHeader("X-API-KEY", c.cfg.SerperAPIKey). + SetHeader("Content-Type", "application/json"). + SetBody(body). + SetResult(&res). + Post(serperScrapeEndpoint) + if err != nil { + return nil, fmt.Errorf("failed to query Serper scrape API: %w", err) + } + + if resp.IsError() { + return nil, fmt.Errorf("Serper scrape API error (status %d): %s", resp.StatusCode(), resp.String()) + } + + return &res, nil + }) + + // Update circuit breaker + c.serperCB.recordResult("serper_scrape", err) + if err != nil { - return nil, fmt.Errorf("failed to query Serper scrape API: %w", err) + return nil, err } - - if resp.IsError() { - return nil, fmt.Errorf("Serper scrape API error (status %d): %s", resp.StatusCode(), resp.String()) + + // Validate response (minimum 50 chars for meaningful content) + if validationErr := ValidateFetchResponse(result, 50); validationErr != nil { + log.Warn().Err(validationErr).Msg("serper scrape returned invalid response") + return EnrichEmptyFetch(result, query.Url, "validation_failed"), nil } - return &result, nil + return result, nil } func (c *SearchClient) fetchFallback(ctx context.Context, query domainsearch.FetchWebpageRequest) (*domainsearch.FetchWebpageResponse, error) { - resp, err := c.fallbackClient.R(). - SetContext(ctx). - SetHeader("User-Agent", "Jan-MCP-Tools-Fallback/1.0"). - Get(query.Url) - if err != nil { - return nil, fmt.Errorf("fallback fetch failed: %w", err) - } - if resp.IsError() { - return nil, fmt.Errorf("fallback fetch HTTP %d: %s", resp.StatusCode(), resp.Status()) - } + // Retry fallback fetch with shorter retry config + shortRetry := c.retryConfig + shortRetry.MaxAttempts = 2 + + result, err := WithRetry(ctx, shortRetry, "fallback_fetch", func() (*domainsearch.FetchWebpageResponse, error) { + resp, err := c.fallbackClient.R(). + SetContext(ctx). + SetHeader("User-Agent", "Jan-MCP-Tools-Fallback/1.0"). + Get(query.Url) + if err != nil { + return nil, fmt.Errorf("fallback fetch failed: %w", err) + } + if resp.IsError() { + return nil, fmt.Errorf("fallback fetch HTTP %d: %s", resp.StatusCode(), resp.Status()) + } - bodyBytes := resp.Body() - text := extractVisibleText(bodyBytes) - if text == "" { - text = string(bodyBytes) - } + bodyBytes := resp.Body() + text := extractVisibleText(bodyBytes) + if text == "" { + text = string(bodyBytes) + } - metadata := map[string]any{ - "source": query.Url, - "contentType": resp.Header().Get("Content-Type"), - "fallback_mode": true, - } + metadata := map[string]any{ + "source": query.Url, + "contentType": resp.Header().Get("Content-Type"), + "fallback_mode": true, + } - return &domainsearch.FetchWebpageResponse{ - Text: text, - Metadata: metadata, - }, nil + return &domainsearch.FetchWebpageResponse{ + Text: text, + Metadata: metadata, + }, nil + }) + + if err != nil { + return nil, err + } + + // Validate response + if validationErr := ValidateFetchResponse(result, 50); validationErr != nil { + log.Warn().Err(validationErr).Msg("fallback fetch returned invalid response") + return EnrichEmptyFetch(result, query.Url, "validation_failed"), nil + } + + return result, nil } func (c *SearchClient) searchViaDuckDuckGo(ctx context.Context, query domainsearch.SearchRequest, reason string) (*domainsearch.SearchResponse, error) { diff --git a/services/mcp-tools/internal/infrastructure/search/retry.go b/services/mcp-tools/internal/infrastructure/search/retry.go new file mode 100644 index 00000000..abbc9dd3 --- /dev/null +++ b/services/mcp-tools/internal/infrastructure/search/retry.go @@ -0,0 +1,129 @@ +package search + +import ( + "context" + "fmt" + "math" + "strings" + "time" + + "github.com/rs/zerolog/log" +) + +// RetryConfig defines retry behavior for search operations +type RetryConfig struct { + MaxAttempts int + InitialDelay time.Duration + MaxDelay time.Duration + BackoffFactor float64 + RetryableErrors []string +} + +// DefaultRetryConfig returns sensible defaults for retry behavior +func DefaultRetryConfig() RetryConfig { + return RetryConfig{ + MaxAttempts: 3, + InitialDelay: 500 * time.Millisecond, + MaxDelay: 10 * time.Second, + BackoffFactor: 2.0, + RetryableErrors: []string{ + "timeout", + "connection refused", + "temporary failure", + "429", // Rate limit + "500", // Internal server error + "502", // Bad gateway + "503", // Service unavailable + "504", // Gateway timeout + }, + } +} + +// RetryableFunc is a function that can be retried +type RetryableFunc[T any] func() (*T, error) + +// WithRetry executes a function with exponential backoff retry logic +func WithRetry[T any](ctx context.Context, cfg RetryConfig, operation string, fn RetryableFunc[T]) (*T, error) { + var lastErr error + + for attempt := 1; attempt <= cfg.MaxAttempts; attempt++ { + result, err := fn() + if err == nil { + if attempt > 1 { + log.Info(). + Str("operation", operation). + Int("attempt", attempt). + Msg("operation succeeded after retry") + } + return result, nil + } + + lastErr = err + + // Check if error is retryable + if !isRetryable(err, cfg.RetryableErrors) { + log.Debug(). + Err(err). + Str("operation", operation). + Int("attempt", attempt). + Msg("non-retryable error, aborting") + return nil, err + } + + // Don't sleep after last attempt + if attempt == cfg.MaxAttempts { + break + } + + // Calculate backoff delay + delay := calculateBackoff(attempt, cfg.InitialDelay, cfg.MaxDelay, cfg.BackoffFactor) + + log.Warn(). + Err(err). + Str("operation", operation). + Int("attempt", attempt). + Int("max_attempts", cfg.MaxAttempts). + Dur("retry_delay", delay). + Msg("retrying operation after error") + + // Wait with context cancellation support + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(delay): + // Continue to next attempt + } + } + + return nil, fmt.Errorf("operation failed after %d attempts: %w", cfg.MaxAttempts, lastErr) +} + +// calculateBackoff computes exponential backoff delay with jitter +func calculateBackoff(attempt int, initial, max time.Duration, factor float64) time.Duration { + backoff := float64(initial) * math.Pow(factor, float64(attempt-1)) + + if backoff > float64(max) { + backoff = float64(max) + } + + // Add 10% jitter to prevent thundering herd + jitter := backoff * 0.1 * (2.0*float64(time.Now().UnixNano()%100)/100.0 - 1.0) + + return time.Duration(backoff + jitter) +} + +// isRetryable checks if an error should trigger a retry +func isRetryable(err error, retryableErrors []string) bool { + if err == nil { + return false + } + + errStr := err.Error() + for _, pattern := range retryableErrors { + if strings.Contains(strings.ToLower(errStr), strings.ToLower(pattern)) { + return true + } + } + + return false +} diff --git a/services/mcp-tools/internal/infrastructure/search/validation.go b/services/mcp-tools/internal/infrastructure/search/validation.go new file mode 100644 index 00000000..b062d946 --- /dev/null +++ b/services/mcp-tools/internal/infrastructure/search/validation.go @@ -0,0 +1,149 @@ +package search + +import ( + "fmt" + "strings" + + domainsearch "jan-server/services/mcp-tools/internal/domain/search" + + "github.com/rs/zerolog/log" +) + +// ValidationError represents a validation failure +type ValidationError struct { + Field string + Message string +} + +func (e ValidationError) Error() string { + return fmt.Sprintf("validation error on %s: %s", e.Field, e.Message) +} + +// ValidateSearchResponse checks if a search response is valid and has meaningful content +func ValidateSearchResponse(resp *domainsearch.SearchResponse, minResults int) error { + if resp == nil { + return ValidationError{Field: "response", Message: "response is nil"} + } + + if resp.Organic == nil { + log.Warn().Msg("search response has nil organic results") + resp.Organic = []map[string]any{} + } + + if len(resp.Organic) == 0 { + return ValidationError{Field: "organic", Message: "no results returned"} + } + + if minResults > 0 && len(resp.Organic) < minResults { + log.Warn(). + Int("expected_min", minResults). + Int("actual", len(resp.Organic)). + Msg("fewer results than expected, but not failing") + } + + // Validate that results have required fields + validResults := 0 + for idx, result := range resp.Organic { + if result == nil { + log.Warn().Int("index", idx).Msg("nil result in organic array") + continue + } + + // Check for essential fields + hasTitle := hasNonEmptyString(result, "title") + hasLink := hasNonEmptyString(result, "link") + + if !hasTitle || !hasLink { + log.Warn(). + Int("index", idx). + Bool("has_title", hasTitle). + Bool("has_link", hasLink). + Msg("result missing essential fields") + continue + } + + validResults++ + } + + if validResults == 0 { + return ValidationError{Field: "organic", Message: "no valid results with title and link"} + } + + return nil +} + +// ValidateFetchResponse checks if a scrape response has meaningful content +func ValidateFetchResponse(resp *domainsearch.FetchWebpageResponse, minLength int) error { + if resp == nil { + return ValidationError{Field: "response", Message: "response is nil"} + } + + text := strings.TrimSpace(resp.Text) + if text == "" { + return ValidationError{Field: "text", Message: "empty text content"} + } + + if minLength > 0 && len(text) < minLength { + return ValidationError{ + Field: "text", + Message: fmt.Sprintf("text too short: %d chars (min: %d)", len(text), minLength), + } + } + + return nil +} + +// EnrichEmptyResponse adds helpful context when responses are empty +func EnrichEmptyResponse(resp *domainsearch.SearchResponse, query string, reason string) *domainsearch.SearchResponse { + if resp == nil { + resp = &domainsearch.SearchResponse{} + } + + if resp.Organic == nil || len(resp.Organic) == 0 { + resp.Organic = []map[string]any{ + { + "title": fmt.Sprintf("No results found for: %s", query), + "link": fmt.Sprintf("https://google.com/search?q=%s", strings.ReplaceAll(query, " ", "+")), + "snippet": fmt.Sprintf("The search returned no results. Reason: %s. Try refining your query or checking connectivity.", reason), + "source": "empty_fallback", + }, + } + } + + if resp.SearchParameters == nil { + resp.SearchParameters = map[string]any{} + } + resp.SearchParameters["empty_result_reason"] = reason + resp.SearchParameters["has_results"] = len(resp.Organic) > 0 + + return resp +} + +// EnrichEmptyFetch adds helpful context when scrape response is empty +func EnrichEmptyFetch(resp *domainsearch.FetchWebpageResponse, url string, reason string) *domainsearch.FetchWebpageResponse { + if resp == nil { + resp = &domainsearch.FetchWebpageResponse{} + } + + if strings.TrimSpace(resp.Text) == "" { + resp.Text = fmt.Sprintf("Failed to fetch content from %s. Reason: %s", url, reason) + } + + if resp.Metadata == nil { + resp.Metadata = map[string]any{} + } + resp.Metadata["empty_result_reason"] = reason + resp.Metadata["has_content"] = strings.TrimSpace(resp.Text) != "" + + return resp +} + +// hasNonEmptyString checks if a map has a non-empty string value for a key +func hasNonEmptyString(m map[string]any, key string) bool { + if val, ok := m[key]; ok { + if str, ok := val.(string); ok { + return strings.TrimSpace(str) != "" + } + } + return false +} diff --git a/services/mcp-tools/internal/interfaces/httpserver/routes/mcp/serper_mcp.go b/services/mcp-tools/internal/interfaces/httpserver/routes/mcp/serper_mcp.go index 00485b70..d8ae68fb 100644 --- a/services/mcp-tools/internal/interfaces/httpserver/routes/mcp/serper_mcp.go +++ b/services/mcp-tools/internal/interfaces/httpserver/routes/mcp/serper_mcp.go @@ -3,7 +3,6 @@ package mcp import ( "context" "encoding/json" - "fmt" "strings" "time" @@ -25,7 +24,6 @@ type SerperSearchArgs struct { Tbs *string `json:"tbs,omitempty" jsonschema:"description=Time-based search filter ('qdr:h' for past hour, 'qdr:d' for past day, 'qdr:w' for past week, 'qdr:m' for past month, 'qdr:y' for past year)"` Page *int `json:"page,omitempty" jsonschema:"description=Page number of results to return (default: 1)"` Autocorrect *bool `json:"autocorrect,omitempty" jsonschema:"description=Whether to autocorrect spelling in query"` - DomainAllowList []string `json:"domain_allow_list,omitempty" jsonschema:"description=Restrict results to the provided domains, e.g., ['example.com','wikipedia.org']"` LocationHint *string `json:"location_hint,omitempty" jsonschema:"description=Soft location hint (region or timezone) applied when the upstream engine supports it"` OfflineMode *bool `json:"offline_mode,omitempty" jsonschema:"description=Force cached/offline search mode even when live engines are available"` } @@ -134,9 +132,6 @@ func (s *SerperMCP) RegisterTools(server *mcpserver.MCPServer) { autocorrect := req.GetBool("autocorrect", true) searchReq.Autocorrect = &autocorrect - if domains := req.GetStringSlice("domain_allow_list", nil); len(domains) > 0 { - searchReq.DomainAllowList = domains - } if locationHint := req.GetString("location_hint", ""); locationHint != "" { searchReq.LocationHint = &locationHint } @@ -199,104 +194,105 @@ func (s *SerperMCP) RegisterTools(server *mcpserver.MCPServer) { }, ) - if s.vectorStore != nil { - server.AddTool( - mcpgo.NewTool("file_search_index", - mcp.ReflectToMCPOptions( - "Index arbitrary text into the lightweight vector store used for MCP automations.", - FileSearchIndexArgs{}, - )..., - ), - func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) { - if s.vectorStore == nil { - return nil, fmt.Errorf("vector store client is not configured") - } - - docID, err := req.RequireString("document_id") - if err != nil { - return nil, err - } - text, err := req.RequireString("text") - if err != nil { - return nil, err - } - - metadata := extractMapArgument(req.GetArguments(), "metadata") - tags := req.GetStringSlice("tags", nil) - - resp, err := s.vectorStore.IndexDocument(ctx, vectorstore.IndexRequest{ - DocumentID: docID, - Text: text, - Metadata: metadata, - Tags: tags, - }) - if err != nil { - return nil, err - } - - payload := map[string]any{ - "document_id": resp.DocumentID, - "status": resp.Status, - "indexed_at": resp.IndexedAt, - "token_count": resp.TokenCount, - } - jsonBytes, err := json.Marshal(payload) - if err != nil { - return nil, err - } - - return mcpgo.NewToolResultText(string(jsonBytes)), nil - }, - ) - - server.AddTool( - mcpgo.NewTool("file_search_query", - mcp.ReflectToMCPOptions( - "Run a semantic query against documents indexed via file_search_index.", - FileSearchQueryArgs{}, - )..., - ), - func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) { - if s.vectorStore == nil { - return nil, fmt.Errorf("vector store client is not configured") - } - - query, err := req.RequireString("query") - if err != nil { - return nil, err - } - - topK := req.GetInt("top_k", 5) - if topK <= 0 { - topK = 5 - } - if topK > 20 { - topK = 20 - } - docIDs := req.GetStringSlice("document_ids", nil) - - resp, err := s.vectorStore.Query(ctx, vectorstore.QueryRequest{ - Text: query, - TopK: topK, - DocumentIDs: docIDs, - }) - if err != nil { - return nil, err - } - - if resp.TopK == 0 { - resp.TopK = topK - } - - jsonBytes, err := json.Marshal(resp) - if err != nil { - return nil, err - } - - return mcpgo.NewToolResultText(string(jsonBytes)), nil - }, - ) - } + // Disabled: file_search_index and file_search_query tools + // if s.vectorStore != nil { + // server.AddTool( + // mcpgo.NewTool("file_search_index", + // mcp.ReflectToMCPOptions( + // "Index arbitrary text into the lightweight vector store used for MCP automations.", + // FileSearchIndexArgs{}, + // )..., + // ), + // func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) { + // if s.vectorStore == nil { + // return nil, fmt.Errorf("vector store client is not configured") + // } + + // docID, err := req.RequireString("document_id") + // if err != nil { + // return nil, err + // } + // text, err := req.RequireString("text") + // if err != nil { + // return nil, err + // } + + // metadata := extractMapArgument(req.GetArguments(), "metadata") + // tags := req.GetStringSlice("tags", nil) + + // resp, err := s.vectorStore.IndexDocument(ctx, vectorstore.IndexRequest{ + // DocumentID: docID, + // Text: text, + // Metadata: metadata, + // Tags: tags, + // }) + // if err != nil { + // return nil, err + // } + + // payload := map[string]any{ + // "document_id": resp.DocumentID, + // "status": resp.Status, + // "indexed_at": resp.IndexedAt, + // "token_count": resp.TokenCount, + // } + // jsonBytes, err := json.Marshal(payload) + // if err != nil { + // return nil, err + // } + + // return mcpgo.NewToolResultText(string(jsonBytes)), nil + // }, + // ) + + // server.AddTool( + // mcpgo.NewTool("file_search_query", + // mcp.ReflectToMCPOptions( + // "Run a semantic query against documents indexed via file_search_index.", + // FileSearchQueryArgs{}, + // )..., + // ), + // func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) { + // if s.vectorStore == nil { + // return nil, fmt.Errorf("vector store client is not configured") + // } + + // query, err := req.RequireString("query") + // if err != nil { + // return nil, err + // } + + // topK := req.GetInt("top_k", 5) + // if topK <= 0 { + // topK = 5 + // } + // if topK > 20 { + // topK = 20 + // } + // docIDs := req.GetStringSlice("document_ids", nil) + + // resp, err := s.vectorStore.Query(ctx, vectorstore.QueryRequest{ + // Text: query, + // TopK: topK, + // DocumentIDs: docIDs, + // }) + // if err != nil { + // return nil, err + // } + + // if resp.TopK == 0 { + // resp.TopK = topK + // } + + // jsonBytes, err := json.Marshal(resp) + // if err != nil { + // return nil, err + // } + + // return mcpgo.NewToolResultText(string(jsonBytes)), nil + // }, + // ) + // } } func buildSearchPayload(query string, req domainsearch.SearchRequest, resp *domainsearch.SearchResponse) searchToolPayload {