diff --git a/tools/ai-review/src/main.rs b/tools/ai-review/src/main.rs index 6751242..e7669f7 100644 --- a/tools/ai-review/src/main.rs +++ b/tools/ai-review/src/main.rs @@ -34,7 +34,9 @@ async fn main() -> anyhow::Result<()> { .context("GITHUB_REPOSITORY must be in owner/repo format")?; let github_token = std::env::var("GITHUB_TOKEN").context("GITHUB_TOKEN not set")?; - let mistral_key = std::env::var("MISTRAL_API_KEY").context("MISTRAL_API_KEY not set")?; + let mistral_key = require_mistral_api_key( + std::env::var("MISTRAL_API_KEY").context("MISTRAL_API_KEY not set")?, + )?; let clients = Clients { octo: OctocrabBuilder::default() @@ -43,6 +45,7 @@ async fn main() -> anyhow::Result<()> { .context("failed to build Octocrab client")?, http: reqwest::Client::builder() .user_agent("ai-review-bot/0.1") + .timeout(mistral::REQUEST_TIMEOUT) .build()?, github_token, mistral_key, @@ -66,6 +69,13 @@ async fn main() -> anyhow::Result<()> { Ok(()) } +fn require_mistral_api_key(mistral_api_key: String) -> anyhow::Result { + if mistral_api_key.trim().is_empty() { + anyhow::bail!("MISTRAL_API_KEY is empty") + } + Ok(mistral_api_key) +} + async fn run_analysis( mode: &Mode, clients: &Clients, @@ -197,3 +207,14 @@ async fn run_describe( println!("Description generated."); Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn empty_mistral_key_is_rejected_before_api_calls() { + assert!(require_mistral_api_key("valid-key".to_string()).is_ok()); + assert!(require_mistral_api_key(" ".to_string()).is_err()); + } +} diff --git a/tools/ai-review/src/mistral.rs b/tools/ai-review/src/mistral.rs index 50c8cf8..b2f4faa 100644 --- a/tools/ai-review/src/mistral.rs +++ b/tools/ai-review/src/mistral.rs @@ -1,6 +1,7 @@ #![allow(clippy::missing_errors_doc)] use std::fmt::Write as _; +use std::time::Duration; use anyhow::Context; use serde::{Deserialize, Serialize}; @@ -10,6 +11,11 @@ use crate::types::{Agent, Lens, LensVerdict, ReviewResponse, Severity, SynthRepo const API_URL: &str = "https://api.mistral.ai/v1/chat/completions"; const MAX_DIFF_BYTES: usize = 20_000; +const MAX_REQUEST_ATTEMPTS: usize = 3; +const RETRY_DELAY_MILLIS: u64 = 250; + +/// Maximum duration of one Mistral API request. +pub(crate) const REQUEST_TIMEOUT: Duration = Duration::from_secs(30); /// Model used by the four specialist agents in team mode. pub(crate) const TEAM_AGENT_MODEL: &str = "codestral-latest"; @@ -619,34 +625,177 @@ async fn send_request( api_key: &str, request: &ChatRequest, ) -> anyhow::Result { - let resp = client - .post(API_URL) - .bearer_auth(api_key) - .json(request) - .send() - .await - .context("failed to reach Mistral API")?; - - let status = resp.status(); - if !status.is_success() { - let body = resp.text().await.unwrap_or_default(); - anyhow::bail!("Mistral API error {status}: {body}"); + send_request_to(client, api_key, request, API_URL).await +} + +async fn send_request_to( + client: &reqwest::Client, + api_key: &str, + request: &ChatRequest, + endpoint: &str, +) -> anyhow::Result { + for attempt in 1..=MAX_REQUEST_ATTEMPTS { + match client + .post(endpoint) + .bearer_auth(api_key) + .json(request) + .send() + .await + { + Ok(response) if response.status().is_success() => { + let chat: ChatResponse = response + .json() + .await + .context("failed to parse Mistral API response")?; + return chat + .choices + .into_iter() + .next() + .map(|choice| choice.message.content) + .context("Mistral returned no choices"); + } + Ok(response) => { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + if is_retryable_status(status) && attempt < MAX_REQUEST_ATTEMPTS { + eprintln!( + "warning: Mistral API returned {status}; retrying request {attempt}/{MAX_REQUEST_ATTEMPTS}" + ); + wait_before_retry(attempt).await; + continue; + } + anyhow::bail!("Mistral API error {status}: {body}"); + } + Err(error) if attempt < MAX_REQUEST_ATTEMPTS => { + eprintln!( + "warning: Mistral API request failed: {error}; retrying request {attempt}/{MAX_REQUEST_ATTEMPTS}" + ); + wait_before_retry(attempt).await; + } + Err(error) => return Err(error).context("failed to reach Mistral API"), + } } - let chat: ChatResponse = resp - .json() - .await - .context("failed to parse Mistral API response")?; - chat.choices - .into_iter() - .next() - .map(|c| c.message.content) - .context("Mistral returned no choices") + unreachable!("the request loop either returns a response or an error") +} + +fn is_retryable_status(status: reqwest::StatusCode) -> bool { + status == reqwest::StatusCode::TOO_MANY_REQUESTS || status.is_server_error() +} + +async fn wait_before_retry(attempt: usize) { + let delay = Duration::from_millis(RETRY_DELAY_MILLIS * attempt as u64); + tokio::time::sleep(delay).await; } #[cfg(test)] mod tests { use super::*; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + const SUCCESS_RESPONSE: &str = "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: close\r\nContent-Length: 42\r\n\r\n{\"choices\":[{\"message\":{\"content\":\"ok\"}}]}"; + const TOO_MANY_REQUESTS_RESPONSE: &str = + "HTTP/1.1 429 Too Many Requests\r\nConnection: close\r\nContent-Length: 0\r\n\r\n"; + const SERVER_ERROR_RESPONSE: &str = + "HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\nContent-Length: 0\r\n\r\n"; + const UNAUTHORIZED_RESPONSE: &str = + "HTTP/1.1 401 Unauthorized\r\nConnection: close\r\nContent-Length: 0\r\n\r\n"; + + async fn response_server( + responses: Vec<&'static str>, + ) -> (String, tokio::task::JoinHandle) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind test server"); + let address = listener.local_addr().expect("read test server address"); + let server = tokio::spawn(async move { + let mut request_count = 0; + for response in responses { + let Ok(Ok((mut stream, _))) = + tokio::time::timeout(Duration::from_secs(2), listener.accept()).await + else { + break; + }; + let mut request = [0_u8; 4_096]; + let _bytes_read = stream.read(&mut request).await.expect("read test request"); + stream + .write_all(response.as_bytes()) + .await + .expect("write test response"); + stream.shutdown().await.expect("close test response"); + request_count += 1; + } + request_count + }); + (format!("http://{address}"), server) + } + + fn test_request() -> ChatRequest { + ChatRequest { + model: "test".to_string(), + messages: Vec::new(), + response_format: ResponseFormat { + kind: "json_object", + }, + temperature: 0.0, + } + } + + #[tokio::test] + async fn retries_a_transient_response_before_returning_success() { + let (endpoint, server) = response_server(vec![ + SERVER_ERROR_RESPONSE, + TOO_MANY_REQUESTS_RESPONSE, + SUCCESS_RESPONSE, + ]) + .await; + let response = send_request_to(&reqwest::Client::new(), "key", &test_request(), &endpoint) + .await + .expect("retry succeeds"); + + assert_eq!(response, "ok"); + assert_eq!(server.await.expect("test server joins"), 3); + } + + #[tokio::test] + async fn does_not_retry_an_authentication_failure() { + let (endpoint, server) = + response_server(vec![UNAUTHORIZED_RESPONSE, SUCCESS_RESPONSE]).await; + let result = + send_request_to(&reqwest::Client::new(), "key", &test_request(), &endpoint).await; + + assert!(result.is_err()); + assert_eq!(server.await.expect("test server joins"), 1); + } + + #[tokio::test] + async fn stops_after_the_bounded_number_of_transient_failures() { + let (endpoint, server) = response_server(vec![ + SERVER_ERROR_RESPONSE, + SERVER_ERROR_RESPONSE, + SERVER_ERROR_RESPONSE, + ]) + .await; + let result = + send_request_to(&reqwest::Client::new(), "key", &test_request(), &endpoint).await; + + assert!(result.is_err()); + assert_eq!( + server.await.expect("test server joins"), + MAX_REQUEST_ATTEMPTS + ); + } + + #[test] + fn retries_transient_responses_but_not_authentication_failures() { + assert!(is_retryable_status(reqwest::StatusCode::TOO_MANY_REQUESTS)); + assert!(is_retryable_status( + reqwest::StatusCode::INTERNAL_SERVER_ERROR + )); + assert!(!is_retryable_status(reqwest::StatusCode::UNAUTHORIZED)); + assert!(!is_retryable_status(reqwest::StatusCode::BAD_REQUEST)); + } #[test] fn synthesis_message_keeps_the_complete_bounded_batch() { diff --git a/tools/ai-review/src/team.rs b/tools/ai-review/src/team.rs index d9ffca7..7ee0a55 100644 --- a/tools/ai-review/src/team.rs +++ b/tools/ai-review/src/team.rs @@ -35,17 +35,7 @@ pub async fn run_team( println!("Fetching PR #{pr_number} diff for team review…"); let mut ctx = github::fetch_diff_context(&clients.octo, owner, repo, pr_number).await?; if ctx.full.trim().is_empty() { - if ctx.coverage_gaps.is_empty() { - println!("Empty diff: nothing to review."); - } else { - let body = review::render_incomplete_team_comment( - "no textual patch was available for review", - &ctx.coverage_gaps, - ); - github::upsert_global_comment(&clients.octo, owner, repo, pr_number, &body, MARKER) - .await?; - } - return Ok(()); + return handle_empty_review_input(clients, owner, repo, pr_number, &ctx).await; } let cache_state = load_team_cache(clients, owner, repo, pr_number, &ctx).await?; @@ -63,25 +53,30 @@ pub async fn run_team( } else { "all specialist agents failed" }; - let body = review::render_incomplete_team_comment(reason, &ctx.coverage_gaps); - github::upsert_global_comment(&clients.octo, owner, repo, pr_number, &body, MARKER).await?; - return Ok(()); + return publish_incomplete_review( + clients, + owner, + repo, + pr_number, + reason, + &ctx.coverage_gaps, + ) + .await; } - let raw_count: usize = batch_runs - .iter() - .flat_map(|run| &run.reports) - .map(|(_, report)| report.findings.len()) - .sum(); + let raw_count = raw_finding_count(&batch_runs); let (synth, synthesis_gaps) = synthesize_batches(clients, &batch_runs).await; ctx.coverage_gaps.extend(synthesis_gaps); let Some(mut synth) = synth else { - let body = review::render_incomplete_team_comment( + return publish_incomplete_review( + clients, + owner, + repo, + pr_number, "every batch synthesis failed", &ctx.coverage_gaps, - ); - github::upsert_global_comment(&clients.octo, owner, repo, pr_number, &body, MARKER).await?; - return Ok(()); + ) + .await; }; let mut findings = std::mem::take(&mut synth.findings); @@ -102,18 +97,13 @@ pub async fn run_team( "Verifying {} finding(s) with the 3-lens vote…", findings.len() ); - let VerificationResults { - verdicts, - cacheable, - } = verify_findings(clients, &ctx, &findings, cache_state.reusable()).await; - let scored: Vec<(SynthFinding, FindingVerdict)> = findings.into_iter().zip(verdicts).collect(); - let verdict = compute_verdict(&scored, capped > 0 || !ctx.coverage_gaps.is_empty()); + let verification = verify_findings(clients, &ctx, &findings, cache_state.reusable()).await; + let scored: Vec<(SynthFinding, FindingVerdict)> = + findings.into_iter().zip(verification.verdicts).collect(); + let has_incomplete_coverage = capped > 0 || !ctx.coverage_gaps.is_empty(); + let verdict = compute_verdict(&scored, has_incomplete_coverage); - let model = format!( - "{} + {}", - mistral::TEAM_AGENT_MODEL, - mistral::TEAM_SYNTH_MODEL - ); + let model = team_model_label(); let view = review::TeamCommentView { executive_summary: &synth.executive_summary, executive_summary_fr: &synth.executive_summary_fr, @@ -131,17 +121,83 @@ pub async fn run_team( coverage_gaps: &ctx.coverage_gaps, }; let mut body = review::render_team_comment(&view); - append_team_cache(&mut body, &cache_state, &ctx, &scored, cacheable)?; + append_team_cache( + &mut body, + &cache_state, + &ctx, + &scored, + verification.cacheable, + has_incomplete_coverage, + )?; println!("Upserting team comment…"); github::upsert_global_comment(&clients.octo, owner, repo, pr_number, &body, MARKER).await?; post_confirmed_inline(clients, owner, repo, pr_number, &head_sha, &ctx, &scored).await?; + ensure_complete_review(has_incomplete_coverage)?; println!("Team review complete."); Ok(()) } +fn team_model_label() -> String { + format!( + "{} + {}", + mistral::TEAM_AGENT_MODEL, + mistral::TEAM_SYNTH_MODEL + ) +} + +async fn handle_empty_review_input( + clients: &Clients, + owner: &str, + repo: &str, + pr_number: u64, + ctx: &github::DiffContext, +) -> anyhow::Result<()> { + if ctx.coverage_gaps.is_empty() { + println!("Empty diff: nothing to review."); + return Ok(()); + } + publish_incomplete_review( + clients, + owner, + repo, + pr_number, + "no textual patch was available for review", + &ctx.coverage_gaps, + ) + .await +} + +fn ensure_complete_review(has_incomplete_coverage: bool) -> anyhow::Result<()> { + if has_incomplete_coverage { + anyhow::bail!("team review is incomplete") + } + Ok(()) +} + +async fn publish_incomplete_review( + clients: &Clients, + owner: &str, + repo: &str, + pr_number: u64, + reason: &str, + coverage_gaps: &[CoverageGap], +) -> anyhow::Result<()> { + let body = review::render_incomplete_team_comment(reason, coverage_gaps); + github::upsert_global_comment(&clients.octo, owner, repo, pr_number, &body, MARKER).await?; + ensure_complete_review(true) +} + +fn raw_finding_count(batch_runs: &[BatchRun]) -> usize { + batch_runs + .iter() + .flat_map(|run| &run.reports) + .map(|(_, report)| report.findings.len()) + .sum() +} + struct TeamCacheState { reviewer_revision: Option, trusted_author_configured: bool, @@ -206,8 +262,9 @@ fn append_team_cache( ctx: &github::DiffContext, scored: &[(SynthFinding, FindingVerdict)], cacheable: Vec, + has_incomplete_coverage: bool, ) -> anyhow::Result<()> { - let complete = !has_transient_coverage_gap(ctx) && cacheable.iter().all(|value| *value); + let complete = can_cache_review(has_incomplete_coverage, &cacheable); if let Some(revision) = state .reviewer_revision .as_deref() @@ -230,6 +287,10 @@ fn append_team_cache( Ok(()) } +fn can_cache_review(has_incomplete_coverage: bool, cacheable: &[bool]) -> bool { + !has_incomplete_coverage && cacheable.iter().all(|value| *value) +} + type AgentReports = Vec<(Agent, crate::types::ReviewResponse)>; struct BatchRun { @@ -561,15 +622,6 @@ async fn verify_findings( } } -fn has_transient_coverage_gap(ctx: &github::DiffContext) -> bool { - ctx.coverage_gaps.iter().any(|gap| { - matches!( - gap.kind, - CoverageGapKind::AgentFailed | CoverageGapKind::SynthesisFailed - ) - }) -} - /// Posts inline comments for confirmed, critical, line-located findings. /// The body shows the French message first and the English message second, /// or just the English message when no translation is available. @@ -895,14 +947,15 @@ mod tests { } #[test] - fn transient_model_gaps_prevent_whole_run_caching() { - let mut ctx = diff_ctx("a.rs", ONE_HUNK_PATCH); - assert!(!has_transient_coverage_gap(&ctx)); - ctx.coverage_gaps.push(CoverageGap { - kind: CoverageGapKind::AgentFailed, - file: "a.rs".to_string(), - detail: "temporary failure".to_string(), - }); - assert!(has_transient_coverage_gap(&ctx)); + fn incomplete_review_returns_an_error_after_its_result_is_published() { + assert!(ensure_complete_review(false).is_ok()); + assert!(ensure_complete_review(true).is_err()); + } + + #[test] + fn incomplete_reviews_are_never_cacheable() { + assert!(can_cache_review(false, &[true, true])); + assert!(!can_cache_review(true, &[true, true])); + assert!(!can_cache_review(false, &[true, false])); } }