use std::sync::Arc;
use tokio::task::JoinSet;
use tracing::info;
use crate::collect::collector::CollectionStats;
use crate::collect::errors::Result;
use crate::collect::pr_provider::PrProvider;
use crate::core::db::Database;
use crate::core::models::PullRequest;
pub(super) async fn drain_and_store_pull_requests(
mut set: JoinSet<(String, Result<Vec<PullRequest>>)>,
providers: &[Arc<dyn PrProvider + Send + Sync>],
db: &mut Database,
stats: &mut CollectionStats,
) {
while let Some(joined) = set.join_next().await {
let (provider_name, fetch_result) = match joined {
Ok(t) => t,
Err(e) => {
stats.fail_stage(format!("PR fetch task panicked: {e}"));
continue;
}
};
let prs = match fetch_result {
Ok(prs) => prs,
Err(e) => {
stats.fail_stage(format!("{provider_name} PR fetch failed: {e}"));
continue;
}
};
let Some(provider) = providers.iter().find(|p| p.name() == provider_name) else {
stats.fail_stage(format!(
"internal: no provider registered for '{provider_name}' when storing PRs"
));
continue;
};
record_blank_head_refs(stats, &provider_name, &prs);
for notice in provider.fetch_notices() {
stats.skip_item(format!("{provider_name}: {notice}"));
}
stats.errors.extend(provider.fetch_faults());
match provider.store_pull_requests(db, &prs) {
Ok(n) => {
info!(provider = %provider_name, prs = n, "stored pull requests");
stats.prs_fetched += n;
}
Err(e) => {
stats.fail_stage(format!("{provider_name} PR store failed: {e}"));
}
}
}
}
fn record_blank_head_refs(stats: &mut CollectionStats, provider: &str, prs: &[PullRequest]) {
let blank = prs
.iter()
.filter(|p| p.head_ref.as_deref() == Some(""))
.count();
if blank > 0 {
stats.skip_item(format!(
"{provider}: {blank} pull request(s) reported an empty source branch; \
no branch ticket key harvested for them"
));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::models::PrState;
use chrono::Utc;
fn pr(head_ref: Option<&str>) -> PullRequest {
PullRequest {
id: 0,
pr_number: 1,
repository: "acme/widgets".into(),
title: "T".into(),
author: "ada".into(),
state: PrState::Merged,
created_at: Utc::now(),
merged_at: None,
commit_shas: "[]".into(),
fetched_at: "2026-01-01T00:00:00Z".into(),
head_ref: head_ref.map(str::to_string),
body_ticket_id: None,
}
}
#[test]
fn blank_head_ref_is_recorded_as_a_skipped_item() {
let mut stats = CollectionStats::default();
record_blank_head_refs(
&mut stats,
"github",
&[pr(Some("")), pr(Some("feature/PROJ-1")), pr(Some(""))],
);
assert_eq!(
stats.errors.len(),
1,
"one aggregated fault, not one per PR"
);
assert!(
stats.stage_failures().is_empty(),
"a payload anomaly must not reach the exit code"
);
let msg = stats.errors[0].message.clone();
assert!(msg.contains("github"), "{msg}");
assert!(msg.contains('2'), "the count must be named: {msg}");
}
#[test]
fn absent_head_ref_is_not_a_fault() {
let mut stats = CollectionStats::default();
record_blank_head_refs(
&mut stats,
"bitbucket",
&[pr(None), pr(None), pr(Some("feature/PROJ-1"))],
);
assert!(stats.errors.is_empty());
}
fn pulls_json(first: usize, count: usize) -> serde_json::Value {
(first..first + count)
.map(|n| {
serde_json::json!({
"number": n, "title": "T", "user": { "login": "ada" },
"state": "open", "created_at": "2026-01-01T00:00:00Z",
"merged_at": null
})
})
.collect()
}
fn denied(status: u16) -> wiremock::ResponseTemplate {
wiremock::ResponseTemplate::new(status)
.set_body_json(serde_json::json!({ "message": "denied or missing" }))
}
async fn drain_github_repos(server: &wiremock::MockServer, slugs: &[&str]) -> CollectionStats {
use crate::collect::github::GitHubClient;
use crate::core::config::GithubConfig;
let cfg = GithubConfig {
token: None,
org: None,
orgs: vec![],
repo: None,
fetch_prs: true,
fetch_pr_reviews: false,
review_fetch_concurrency: 1,
ticket_regex: None,
fetch_on_reference: false,
work_items_unavailable: None,
};
let repos = slugs
.iter()
.filter_map(|slug| slug.split_once('/'))
.map(|(o, r)| (o.to_string(), r.to_string()))
.collect();
let client = GitHubClient::new_for_prs(&cfg, repos)
.expect("client builds")
.with_api_base(server.uri());
let provider: Arc<dyn PrProvider + Send + Sync> = Arc::new(client);
let providers = vec![Arc::clone(&provider)];
let mut set = JoinSet::new();
set.spawn(async move { ("github".to_string(), provider.fetch_pull_requests().await) });
let mut db = Database::open_in_memory().expect("open db");
let mut stats = CollectionStats::default();
drain_and_store_pull_requests(set, &providers, &mut db, &mut stats).await;
stats
}
async fn drain_github_against(answers: &[(&str, u16)]) -> CollectionStats {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
for (i, (slug, status)) in answers.iter().enumerate() {
let reply = if *status == 200 {
ResponseTemplate::new(200).set_body_json(pulls_json(i + 1, 1))
} else {
denied(*status)
};
Mock::given(method("GET"))
.and(path(format!("/repos/{slug}/pulls")))
.respond_with(reply)
.mount(&server)
.await;
}
let slugs: Vec<&str> = answers.iter().map(|(slug, _)| *slug).collect();
drain_github_repos(&server, &slugs).await
}
async fn drain_github_page2_fails(page2_status: u16) -> CollectionStats {
use crate::collect::github::client::PAGE_SIZE;
use wiremock::matchers::{method, path, query_param};
use wiremock::{Mock, MockServer, ResponseTemplate};
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/repos/acme/widgets/pulls"))
.and(query_param("page", "1"))
.respond_with(
ResponseTemplate::new(200).set_body_json(pulls_json(1, PAGE_SIZE as usize)),
)
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/repos/acme/widgets/pulls"))
.and(query_param("page", "2"))
.respond_with(denied(page2_status))
.mount(&server)
.await;
drain_github_repos(&server, &["acme/widgets"]).await
}
#[tokio::test]
async fn a_pr_list_404_is_a_counted_warning_not_a_stage_failure() {
let stats = drain_github_against(&[
("acme/widgets", 200),
("acme/gone", 404),
("acme/zeta", 200),
])
.await;
assert_eq!(stats.prs_fetched, 2, "the repos around the 404 still store");
assert!(
stats.stage_failures().is_empty(),
"a 404 must not fail the PR stage; got: {:?}",
stats.errors
);
assert_eq!(
stats.errors.len(),
1,
"a 404 must be one counted warning, not silent; got: {:?}",
stats.errors
);
let msg = &stats.errors[0].message;
for needle in ["1 of 3", "HTTP 404", "acme/gone"] {
assert!(msg.contains(needle), "missing `{needle}` in: {msg}");
}
}
#[tokio::test]
async fn a_pr_list_403_fails_the_stage_once() {
let stats = drain_github_against(&[
("acme/widgets", 200),
("acme/secret", 403),
("acme/zeta", 200),
])
.await;
assert_eq!(stats.prs_fetched, 2, "the repos around the 403 still store");
let failures = stats.stage_failures();
assert_eq!(
failures.len(),
1,
"a 403 must fail the PR stage exactly once; got: {:?}",
stats.errors
);
assert_eq!(stats.errors.len(), 1, "no extra faults: {:?}", stats.errors);
let msg = &failures[0].message;
for needle in ["1 of 3", "HTTP 403: 1", "acme/secret (HTTP 403)"] {
assert!(msg.contains(needle), "missing `{needle}` in: {msg}");
}
}
#[tokio::test]
async fn a_page_2_404_keeps_page_1_and_is_a_counted_warning() {
use crate::collect::github::client::PAGE_SIZE;
let stats = drain_github_page2_fails(404).await;
assert_eq!(
stats.prs_fetched, PAGE_SIZE as usize,
"page 1's pull requests must be stored"
);
assert!(
stats.stage_failures().is_empty(),
"a page-2 404 must not fail the PR stage; got: {:?}",
stats.errors
);
assert_eq!(
stats.errors.len(),
1,
"a partial fetch must be one counted warning; got: {:?}",
stats.errors
);
let msg = &stats.errors[0].message;
for needle in [
"1 of 1",
"HTTP 404",
"acme/widgets (pull requests after page 1 not collected)",
] {
assert!(msg.contains(needle), "missing `{needle}` in: {msg}");
}
}
#[tokio::test]
async fn a_page_2_403_keeps_page_1_and_fails_the_stage_once() {
use crate::collect::github::client::PAGE_SIZE;
let stats = drain_github_page2_fails(403).await;
assert_eq!(
stats.prs_fetched, PAGE_SIZE as usize,
"page 1's pull requests must be stored"
);
let failures = stats.stage_failures();
assert_eq!(
failures.len(),
1,
"a page-2 403 must fail the PR stage exactly once; got: {:?}",
stats.errors
);
assert_eq!(stats.errors.len(), 1, "no extra faults: {:?}", stats.errors);
let msg = &failures[0].message;
for needle in [
"1 of 1",
"HTTP 403: 1",
"acme/widgets (pull requests after page 1 not collected) (HTTP 403)",
] {
assert!(msg.contains(needle), "missing `{needle}` in: {msg}");
}
}
}