diff --git a/internal/ai/cost_wrapper.go b/internal/ai/cost_wrapper.go index 8a3d26f..3994e3a 100644 --- a/internal/ai/cost_wrapper.go +++ b/internal/ai/cost_wrapper.go @@ -80,6 +80,14 @@ func (w *CostAwareWrapper) SetSkipConfirmation(skip bool) { w.skipConfirmation = skip } +// InvalidateCache removes the cached response for the given prompt, so a +// truncated or malformed response (network-successful but unusable) doesn't +// keep getting served back on retry until the TTL expires. +func (w *CostAwareWrapper) InvalidateCache(prompt string) error { + contentHash := w.cache.GenerateHash(w.provider.GetProviderName() + w.provider.GetModelName() + prompt) + return w.cache.Delete(contentHash) +} + // WrapGenerate wraps any generation function with tracking func (w *CostAwareWrapper) WrapGenerate( ctx context.Context, diff --git a/internal/ai/gemini/commit_summarizer_service.go b/internal/ai/gemini/commit_summarizer_service.go index def11c5..2e1b0ac 100644 --- a/internal/ai/gemini/commit_summarizer_service.go +++ b/internal/ai/gemini/commit_summarizer_service.go @@ -223,63 +223,89 @@ func (s *GeminiCommitSummarizer) GenerateSuggestions(ctx context.Context, info m "prompt_length", len(prompt), "language", s.config.Language) - resp, usage, err := s.wrapper.WrapGenerate(ctx, "suggest-commits", prompt, s.generateFn) - if err != nil { - log.Error("failed to generate suggestions", - "error", err) - return nil, err - } + var usage *models.TokenUsage + var suggestions []models.CommitSuggestion + + const maxAttempts = 2 + for attempt := 1; attempt <= maxAttempts; attempt++ { + resp, respUsage, err := s.wrapper.WrapGenerate(ctx, "suggest-commits", prompt, s.generateFn) + if err != nil { + log.Error("failed to generate suggestions", + "error", err) + return nil, err + } - var responseText string - if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { - log.Debug("formatResponse received GenerateContentResponse", - "candidates_count", len(geminiResp.Candidates)) - responseText = formatResponse(geminiResp) - if len(responseText) > 0 { - preview := responseText - if len(responseText) > 100 { - preview = responseText[:100] + var responseText string + if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { + log.Debug("formatResponse received GenerateContentResponse", + "candidates_count", len(geminiResp.Candidates)) + responseText = formatResponse(geminiResp) + if len(responseText) > 0 { + preview := responseText + if len(responseText) > 100 { + preview = responseText[:100] + } + log.Debug("formatResponse result", + "response_length", len(responseText), + "response_preview", preview) + } else { + log.Debug("formatResponse result empty") } - log.Debug("formatResponse result", - "response_length", len(responseText), - "response_preview", preview) + } else if str, ok := resp.(string); ok { + responseText = str + log.Debug("received string response", "length", len(str)) + } else if respMap, ok := resp.(map[string]interface{}); ok { + log.Debug("received map response from cache, extracting text") + responseText = extractTextFromMap(respMap) + log.Debug("extracted text from map", "length", len(responseText)) } else { - log.Debug("formatResponse result empty") + log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) } - } else if str, ok := resp.(string); ok { - responseText = str - log.Debug("received string response", "length", len(str)) - } else if respMap, ok := resp.(map[string]interface{}); ok { - log.Debug("received map response from cache, extracting text") - responseText = extractTextFromMap(respMap) - log.Debug("extracted text from map", "length", len(responseText)) - } else { - log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) - } - if responseText == "" { - return nil, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "empty response from AI"). - WithContext("operation", "generate commit suggestions") - } + if responseText == "" { + if attempt < maxAttempts { + log.Warn("empty response from AI, invalidating cache and retrying", "attempt", attempt) + _ = s.wrapper.InvalidateCache(prompt) + continue + } + return nil, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "empty response from AI"). + WithContext("operation", "generate commit suggestions") + } - suggestions, err := s.parseSuggestionsJSON(responseText) - if err != nil { - respLen := len(responseText) - preview := responseText - if respLen > 500 { - preview = responseText[:500] + "..." + parsed, err := s.parseSuggestionsJSON(responseText) + if err != nil { + if attempt < maxAttempts { + log.Warn("failed to parse suggestions JSON, invalidating cache and retrying", + "error", err, "attempt", attempt) + _ = s.wrapper.InvalidateCache(prompt) + continue + } + respLen := len(responseText) + preview := responseText + if respLen > 500 { + preview = responseText[:500] + "..." + } + return nil, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "failed to parse JSON"). + WithContext("response_length", respLen). + WithContext("preview", preview). + WithError(err) + } + if len(parsed) == 0 { + if attempt < maxAttempts { + log.Warn("AI generated no suggestions, invalidating cache and retrying", "attempt", attempt) + _ = s.wrapper.InvalidateCache(prompt) + continue + } + log.Warn("AI generated no suggestions") + return nil, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "AI generated no suggestions") } - return nil, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "failed to parse JSON"). - WithContext("response_length", respLen). - WithContext("preview", preview). - WithError(err) - } - if len(suggestions) == 0 { - log.Warn("AI generated no suggestions") - return nil, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "AI generated no suggestions") + + usage = respUsage + suggestions = parsed + break } for i := range suggestions { suggestions[i].Usage = usage diff --git a/internal/ai/gemini/issue_content_generator.go b/internal/ai/gemini/issue_content_generator.go index 98f55ac..99f4909 100644 --- a/internal/ai/gemini/issue_content_generator.go +++ b/internal/ai/gemini/issue_content_generator.go @@ -142,43 +142,56 @@ func (s *GeminiIssueContentGenerator) GenerateIssueContent(ctx context.Context, log.Debug("calling gemini API for issue content", "prompt_length", len(prompt)) - resp, usage, err := s.wrapper.WrapGenerate(ctx, "generate-issue", prompt, s.generateFn) - if err != nil { - log.Error("failed to generate issue content", - "error", err) - return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error generating issue content", err) - } - + var usage *models.TokenUsage var responseText string - if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { - log.Debug("formatResponse received GenerateContentResponse", - "candidates_count", len(geminiResp.Candidates)) - responseText = formatResponse(geminiResp) - if len(responseText) > 0 { - preview := responseText - if len(responseText) > 100 { - preview = responseText[:100] + + const maxAttempts = 2 + for attempt := 1; attempt <= maxAttempts; attempt++ { + resp, respUsage, err := s.wrapper.WrapGenerate(ctx, "generate-issue", prompt, s.generateFn) + if err != nil { + log.Error("failed to generate issue content", + "error", err) + return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error generating issue content", err) + } + + if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { + log.Debug("formatResponse received GenerateContentResponse", + "candidates_count", len(geminiResp.Candidates)) + responseText = formatResponse(geminiResp) + if len(responseText) > 0 { + preview := responseText + if len(responseText) > 100 { + preview = responseText[:100] + } + log.Debug("formatResponse result", + "response_length", len(responseText), + "response_preview", preview) + } else { + log.Debug("formatResponse result empty") } - log.Debug("formatResponse result", - "response_length", len(responseText), - "response_preview", preview) + } else if str, ok := resp.(string); ok { + responseText = str + log.Debug("received string response", "length", len(str)) + } else if respMap, ok := resp.(map[string]interface{}); ok { + log.Debug("received map response from cache, extracting text") + responseText = extractTextFromMap(respMap) + log.Debug("extracted text from map", "length", len(responseText)) } else { - log.Debug("formatResponse result empty") + log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) } - } else if str, ok := resp.(string); ok { - responseText = str - log.Debug("received string response", "length", len(str)) - } else if respMap, ok := resp.(map[string]interface{}); ok { - log.Debug("received map response from cache, extracting text") - responseText = extractTextFromMap(respMap) - log.Debug("extracted text from map", "length", len(responseText)) - } else { - log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) - } - if responseText == "" { - log.Error("empty response from gemini AI after format") - return nil, domainErrors.NewAppError(domainErrors.TypeAI, "empty response from AI", nil) + if responseText == "" { + if attempt < maxAttempts { + log.Warn("empty response from AI, invalidating cache and retrying", "attempt", attempt) + _ = s.wrapper.InvalidateCache(prompt) + continue + } + log.Error("empty response from gemini AI after format") + return nil, domainErrors.NewAppError(domainErrors.TypeAI, "empty response from AI", nil) + } + + usage = respUsage + break } log.Debug("gemini response received", diff --git a/internal/ai/gemini/pull_requests_summarizer_service.go b/internal/ai/gemini/pull_requests_summarizer_service.go index 17d8056..baf5f4f 100644 --- a/internal/ai/gemini/pull_requests_summarizer_service.go +++ b/internal/ai/gemini/pull_requests_summarizer_service.go @@ -138,60 +138,86 @@ func (gps *GeminiPRSummarizer) GeneratePRSummary(ctx context.Context, prContent log.Debug("calling gemini API for PR summary", "prompt_length", len(prompt)) - resp, usage, err := gps.wrapper.WrapGenerate(ctx, "summarize-pr", prompt, gps.generateFn) - if err != nil { - log.Error("failed to generate PR summary", - "error", err) - return models.PRSummary{}, err - } + var usage *models.TokenUsage + var jsonSummary PRSummaryJSON - var responseText string - if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { - log.Debug("formatResponse received GenerateContentResponse", - "candidates_count", len(geminiResp.Candidates)) - responseText = formatResponse(geminiResp) - } else if str, ok := resp.(string); ok { - responseText = str - log.Debug("received string response", "length", len(str)) - } else if respMap, ok := resp.(map[string]interface{}); ok { - log.Debug("received map response from cache, extracting text") - responseText = extractTextFromMap(respMap) - log.Debug("extracted text from map", "length", len(responseText)) - } else { - log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) - } + const maxAttempts = 2 + for attempt := 1; attempt <= maxAttempts; attempt++ { + resp, respUsage, err := gps.wrapper.WrapGenerate(ctx, "summarize-pr", prompt, gps.generateFn) + if err != nil { + log.Error("failed to generate PR summary", + "error", err) + return models.PRSummary{}, err + } - if responseText == "" { - return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "empty response from AI"). - WithContext("operation", "summarize PR") - } + var responseText string + if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { + log.Debug("formatResponse received GenerateContentResponse", + "candidates_count", len(geminiResp.Candidates)) + responseText = formatResponse(geminiResp) + } else if str, ok := resp.(string); ok { + responseText = str + log.Debug("received string response", "length", len(str)) + } else if respMap, ok := resp.(map[string]interface{}); ok { + log.Debug("received map response from cache, extracting text") + responseText = extractTextFromMap(respMap) + log.Debug("extracted text from map", "length", len(responseText)) + } else { + log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) + } - var jsonSummary PRSummaryJSON - if err := json.Unmarshal([]byte(responseText), &jsonSummary); err != nil { - respLen := len(responseText) - preview := responseText - if respLen > 500 { - preview = responseText[:500] + "..." + if responseText == "" { + if attempt < maxAttempts { + log.Warn("empty response from AI, invalidating cache and retrying", "attempt", attempt) + _ = gps.wrapper.InvalidateCache(prompt) + continue + } + return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "empty response from AI"). + WithContext("operation", "summarize PR") } - return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "failed to parse JSON"). - WithContext("response_length", respLen). - WithContext("preview", preview). - WithError(err) - } - if strings.TrimSpace(jsonSummary.Title) == "" { - respLen := len(responseText) - preview := responseText - if respLen > 500 { - preview = responseText[:500] + "..." + + var parsed PRSummaryJSON + if err := json.Unmarshal([]byte(responseText), &parsed); err != nil { + if attempt < maxAttempts { + log.Warn("failed to parse PR summary JSON, invalidating cache and retrying", + "error", err, "attempt", attempt) + _ = gps.wrapper.InvalidateCache(prompt) + continue + } + respLen := len(responseText) + preview := responseText + if respLen > 500 { + preview = responseText[:500] + "..." + } + return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "failed to parse JSON"). + WithContext("response_length", respLen). + WithContext("preview", preview). + WithError(err) + } + if strings.TrimSpace(parsed.Title) == "" { + if attempt < maxAttempts { + log.Warn("AI generated no PR title, invalidating cache and retrying", "attempt", attempt) + _ = gps.wrapper.InvalidateCache(prompt) + continue + } + respLen := len(responseText) + preview := responseText + if respLen > 500 { + preview = responseText[:500] + "..." + } + log.Warn("AI generated no PR title", + "response_length", respLen) + return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "AI generated no PR title"). + WithContext("response_length", respLen). + WithContext("preview", preview) } - log.Warn("AI generated no PR title", - "response_length", respLen) - return models.PRSummary{}, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "AI generated no PR title"). - WithContext("response_length", respLen). - WithContext("preview", preview) + + usage = respUsage + jsonSummary = parsed + break } log.Info("PR summary generated successfully via gemini", diff --git a/internal/ai/gemini/release_generator.go b/internal/ai/gemini/release_generator.go index c7b714d..03c875a 100644 --- a/internal/ai/gemini/release_generator.go +++ b/internal/ai/gemini/release_generator.go @@ -186,54 +186,72 @@ func (g *ReleaseNotesGenerator) GenerateNotes(ctx context.Context, release *mode log.Debug("calling gemini API for release notes", "prompt_length", len(prompt)) - resp, usage, err := g.wrapper.WrapGenerate(ctx, "generate-release", prompt, g.generateFn) - if err != nil { - log.Error("failed to generate release notes", - "error", err, - "version", release.Version) - return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error generating release notes", err) - } + const maxAttempts = 2 + for attempt := 1; attempt <= maxAttempts; attempt++ { + resp, usage, err := g.wrapper.WrapGenerate(ctx, "generate-release", prompt, g.generateFn) + if err != nil { + log.Error("failed to generate release notes", + "error", err, + "version", release.Version) + return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error generating release notes", err) + } - var responseText string - if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { - log.Debug("formatResponse received GenerateContentResponse", - "candidates_count", len(geminiResp.Candidates)) - responseText = formatResponse(geminiResp) - } else if str, ok := resp.(string); ok { - responseText = str - log.Debug("received string response", "length", len(str)) - } else if respMap, ok := resp.(map[string]interface{}); ok { - log.Debug("received map response from cache, extracting text") - responseText = extractTextFromMap(respMap) - log.Debug("extracted text from map", "length", len(responseText)) - } else { - log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) - } + var responseText string + if geminiResp, ok := resp.(*genai.GenerateContentResponse); ok { + log.Debug("formatResponse received GenerateContentResponse", + "candidates_count", len(geminiResp.Candidates)) + responseText = formatResponse(geminiResp) + } else if str, ok := resp.(string); ok { + responseText = str + log.Debug("received string response", "length", len(str)) + } else if respMap, ok := resp.(map[string]interface{}); ok { + log.Debug("received map response from cache, extracting text") + responseText = extractTextFromMap(respMap) + log.Debug("extracted text from map", "length", len(responseText)) + } else { + log.Warn("unexpected response type", "type", fmt.Sprintf("%T", resp)) + } - if responseText == "" { - log.Error("empty response from gemini AI") - return nil, domainErrors.ErrInvalidAIOutput. - WithContext("reason", "empty response from AI"). - WithContext("operation", "generate release notes") - } + if responseText == "" { + if attempt < maxAttempts { + log.Warn("empty response from gemini AI, invalidating cache and retrying", "attempt", attempt) + _ = g.wrapper.InvalidateCache(prompt) + continue + } + log.Error("empty response from gemini AI") + return nil, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "empty response from AI"). + WithContext("operation", "generate release notes") + } - log.Debug("gemini response received", - "response_length", len(responseText)) + log.Debug("gemini response received", + "response_length", len(responseText)) - notes, err := g.parseJSONResponse(responseText, release) - if err != nil { - log.Error("failed to parse release notes response", - "error", err) - return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error parsing AI JSON response", err) - } + notes, err := g.parseJSONResponse(responseText, release) + if err != nil { + if attempt < maxAttempts { + log.Warn("failed to parse release notes response, invalidating cache and retrying", + "error", err, "attempt", attempt) + _ = g.wrapper.InvalidateCache(prompt) + continue + } + log.Error("failed to parse release notes response", + "error", err) + return nil, domainErrors.NewAppError(domainErrors.TypeAI, "error parsing AI JSON response", err) + } - notes.Usage = usage + notes.Usage = usage - log.Info("release notes generated successfully via gemini", - "title", notes.Title, - "highlights_count", len(notes.Highlights)) + log.Info("release notes generated successfully via gemini", + "title", notes.Title, + "highlights_count", len(notes.Highlights)) - return notes, nil + return notes, nil + } + + return nil, domainErrors.ErrInvalidAIOutput. + WithContext("reason", "AI response invalid after retry"). + WithContext("operation", "generate release notes") } func (g *ReleaseNotesGenerator) buildPrompt(release *models.Release) string { diff --git a/internal/ai/gemini/release_generator_test.go b/internal/ai/gemini/release_generator_test.go index 39413b9..704828d 100644 --- a/internal/ai/gemini/release_generator_test.go +++ b/internal/ai/gemini/release_generator_test.go @@ -338,4 +338,53 @@ func TestGenerateNotes(t *testing.T) { assert.Error(t, err) assert.Contains(t, err.Error(), "invalid AI output format") }) + + t.Run("retries once on truncated JSON and succeeds on second attempt", func(t *testing.T) { + // Arrange + callCount := 0 + validJSON := `{"title": "Release v2.0.0", "summary": "Summary", "highlights": ["H1"], "breaking_changes": []}` + generator.generateFn = func(ctx context.Context, mName string, p string) (interface{}, *models.TokenUsage, error) { + callCount++ + text := `{"title": "trunc` // malformed/truncated JSON on the first attempt + if callCount > 1 { + text = validJSON + } + return &genai.GenerateContentResponse{ + Candidates: []*genai.Candidate{ + {Content: &genai.Content{Parts: []*genai.Part{{Text: text}}}}, + }, + UsageMetadata: &genai.GenerateContentResponseUsageMetadata{TotalTokenCount: 100}, + }, &models.TokenUsage{TotalTokens: 100}, nil + } + + // Act + notes, err := generator.GenerateNotes(ctx, &models.Release{Version: "v2.0.0-retry-ok"}) + + // Assert + assert.NoError(t, err) + assert.Equal(t, "Release v2.0.0", notes.Title) + assert.Equal(t, 2, callCount, "should retry exactly once after the truncated response") + }) + + t.Run("gives up after exhausting retries on persistently truncated JSON", func(t *testing.T) { + // Arrange + callCount := 0 + generator.generateFn = func(ctx context.Context, mName string, p string) (interface{}, *models.TokenUsage, error) { + callCount++ + return &genai.GenerateContentResponse{ + Candidates: []*genai.Candidate{ + {Content: &genai.Content{Parts: []*genai.Part{{Text: `{"title": "still trunc`}}}}, + }, + UsageMetadata: &genai.GenerateContentResponseUsageMetadata{TotalTokenCount: 100}, + }, &models.TokenUsage{TotalTokens: 100}, nil + } + + // Act + notes, err := generator.GenerateNotes(ctx, &models.Release{Version: "v2.0.0-retry-fail"}) + + // Assert + assert.Error(t, err) + assert.Nil(t, notes) + assert.Equal(t, 2, callCount, "should stop after the second attempt, not retry forever") + }) } diff --git a/internal/cache/cache.go b/internal/cache/cache.go index 71e1f93..2bf2771 100644 --- a/internal/cache/cache.go +++ b/internal/cache/cache.go @@ -99,6 +99,15 @@ func (c *Cache) Set(hash string, response interface{}) error { return nil } +// Delete removes a single cached response by hash, if present. +func (c *Cache) Delete(hash string) error { + filePath := filepath.Join(c.cacheDir, hash+".json") + if err := os.Remove(filePath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("error deleting cache entry: %w", err) + } + return nil +} + // CleanExpired removes expired cache files func (c *Cache) CleanExpired() error { entries, err := os.ReadDir(c.cacheDir)