use crate::{Config, BranchSpec, Error, Result};
use petgraph::{Graph, Direction};
use petgraph::graph::NodeIndex;
use petgraph::visit::EdgeRef;
use std::collections::{HashMap, HashSet, VecDeque};
use tracing::{debug, info};
pub struct PartyResolver<'a> {
config: &'a Config,
}
#[derive(Debug, Clone)]
pub struct ResolvedParty {
pub name: String,
pub branches: Vec<String>,
pub resolution_order: Vec<String>,
}
impl<'a> PartyResolver<'a> {
pub fn new(config: &'a Config) -> Self {
Self { config }
}
pub fn resolve_party(&self, party_name: &str) -> Result<ResolvedParty> {
info!("Resolving party: {}", party_name);
let _party = self.config.parties.get(party_name)
.ok_or_else(|| Error::party_not_found(party_name))?;
let mut graph = Graph::new();
let mut node_map = HashMap::new();
let mut reverse_map = HashMap::new();
self.collect_all_parties(party_name, &mut node_map, &mut reverse_map, &mut graph)?;
self.add_party_edges(party_name, &node_map, &mut graph)?;
if let Some(cycle) = self.detect_cycles(&graph, &reverse_map) {
return Err(Error::party_cycle(party_name, cycle));
}
let resolution_order = self.topological_sort(&graph, &reverse_map)?;
let branches = self.expand_to_branches(&resolution_order)?;
debug!("Resolved party '{}' to {} branches: {:?}", party_name, branches.len(), branches);
Ok(ResolvedParty {
name: party_name.to_string(),
branches,
resolution_order,
})
}
fn collect_all_parties(
&self,
party_name: &str,
node_map: &mut HashMap<String, NodeIndex>,
reverse_map: &mut HashMap<NodeIndex, String>,
graph: &mut Graph<String, ()>,
) -> Result<()> {
let mut visited = HashSet::new();
let mut queue = VecDeque::new();
queue.push_back(party_name.to_string());
while let Some(current_party) = queue.pop_front() {
if visited.contains(¤t_party) {
continue;
}
visited.insert(current_party.clone());
if !node_map.contains_key(¤t_party) {
let node_idx = graph.add_node(current_party.clone());
node_map.insert(current_party.clone(), node_idx);
reverse_map.insert(node_idx, current_party.clone());
}
let party = self.config.parties.get(¤t_party)
.ok_or_else(|| Error::party_not_found(¤t_party))?;
for member in &party.members {
let spec = BranchSpec::parse(member);
if spec.is_party {
queue.push_back(spec.name);
}
}
}
Ok(())
}
fn add_party_edges(
&self,
party_name: &str,
node_map: &HashMap<String, NodeIndex>,
graph: &mut Graph<String, ()>,
) -> Result<()> {
let mut visited = HashSet::new();
let mut queue = VecDeque::new();
queue.push_back(party_name.to_string());
while let Some(current_party) = queue.pop_front() {
if visited.contains(¤t_party) {
continue;
}
visited.insert(current_party.clone());
let party = self.config.parties.get(¤t_party)
.ok_or_else(|| Error::party_not_found(¤t_party))?;
let current_node = node_map[¤t_party];
for member in &party.members {
let spec = BranchSpec::parse(member);
if spec.is_party {
let dependency_node = node_map[&spec.name];
graph.add_edge(dependency_node, current_node, ());
queue.push_back(spec.name);
}
}
}
Ok(())
}
fn detect_cycles(
&self,
graph: &Graph<String, ()>,
reverse_map: &HashMap<NodeIndex, String>,
) -> Option<String> {
let mut white = HashSet::new(); let mut gray = HashSet::new(); let mut black = HashSet::new();
for node_idx in graph.node_indices() {
white.insert(node_idx);
}
while let Some(&start_node) = white.iter().next() {
if let Some(cycle_path) = self.dfs_cycle_detect(start_node, graph, &mut white, &mut gray, &mut black) {
let cycle_names: Vec<String> = cycle_path.iter()
.map(|&idx| reverse_map[&idx].clone())
.collect();
return Some(cycle_names.join(" -> "));
}
}
None
}
fn dfs_cycle_detect(
&self,
node: NodeIndex,
graph: &Graph<String, ()>,
white: &mut HashSet<NodeIndex>,
gray: &mut HashSet<NodeIndex>,
black: &mut HashSet<NodeIndex>,
) -> Option<Vec<NodeIndex>> {
white.remove(&node);
gray.insert(node);
for edge in graph.edges_directed(node, Direction::Outgoing) {
let neighbor = edge.target();
if gray.contains(&neighbor) {
return Some(vec![node, neighbor]);
}
if white.contains(&neighbor) {
if let Some(mut cycle_path) = self.dfs_cycle_detect(neighbor, graph, white, gray, black) {
cycle_path.insert(0, node);
return Some(cycle_path);
}
}
}
gray.remove(&node);
black.insert(node);
None
}
fn topological_sort(
&self,
graph: &Graph<String, ()>,
reverse_map: &HashMap<NodeIndex, String>,
) -> Result<Vec<String>> {
let mut in_degree = HashMap::new();
let mut queue = VecDeque::new();
let mut result = Vec::new();
for node_idx in graph.node_indices() {
in_degree.insert(node_idx, graph.edges_directed(node_idx, Direction::Incoming).count());
}
for (&node_idx, °ree) in &in_degree {
if degree == 0 {
queue.push_back(node_idx);
}
}
while let Some(node_idx) = queue.pop_front() {
result.push(reverse_map[&node_idx].clone());
for edge in graph.edges_directed(node_idx, Direction::Outgoing) {
let neighbor = edge.target();
if let Some(degree) = in_degree.get_mut(&neighbor) {
*degree -= 1;
if *degree == 0 {
queue.push_back(neighbor);
}
}
}
}
if result.len() != graph.node_count() {
return Err(Error::party_cycle("unknown".to_string(), "cycle detected in topological sort".to_string()));
}
Ok(result)
}
fn expand_to_branches(&self, resolution_order: &[String]) -> Result<Vec<String>> {
let mut branches = Vec::new();
let mut seen_branches = HashSet::new();
for party_name in resolution_order {
let party = self.config.parties.get(party_name)
.ok_or_else(|| Error::party_not_found(party_name))?;
for member in &party.members {
let spec = BranchSpec::parse(member);
if !spec.is_party {
if seen_branches.insert(spec.name.clone()) {
branches.push(spec.name);
}
}
}
}
Ok(branches)
}
pub fn get_dependencies(&self, party_name: &str) -> Result<Vec<String>> {
let resolved = self.resolve_party(party_name)?;
Ok(resolved.resolution_order.into_iter().filter(|name| name != party_name).collect())
}
pub fn validate_party_references(&self, party_name: &str) -> Result<()> {
let party = self.config.parties.get(party_name)
.ok_or_else(|| Error::party_not_found(party_name))?;
for member in &party.members {
let spec = BranchSpec::parse(member);
if spec.is_party && !self.config.parties.contains_key(&spec.name) {
return Err(Error::party_not_found(&spec.name));
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{Config, Party};
use std::collections::HashMap;
fn create_test_config() -> Config {
let mut parties = HashMap::new();
parties.insert("qa".to_string(), Party {
members: vec!["feature/a".to_string(), "feature/b".to_string(), "@payments".to_string()],
..Party::default()
});
parties.insert("payments".to_string(), Party {
members: vec!["feature/pay-ux".to_string(), "feature/pay-api".to_string()],
..Party::default()
});
parties.insert("frontend".to_string(), Party {
members: vec!["feature/ui-update".to_string(), "@payments".to_string()],
..Party::default()
});
Config {
base_branch: "main".to_string(),
parties,
auto_update: crate::AutoUpdateConfig::default(),
}
}
#[test]
fn test_simple_party_resolution() {
let config = create_test_config();
let resolver = PartyResolver::new(&config);
let resolved = resolver.resolve_party("payments").unwrap();
assert_eq!(resolved.name, "payments");
assert_eq!(resolved.branches, vec!["feature/pay-ux", "feature/pay-api"]);
}
#[test]
fn test_nested_party_resolution() {
let config = create_test_config();
let resolver = PartyResolver::new(&config);
let resolved = resolver.resolve_party("qa").unwrap();
assert_eq!(resolved.name, "qa");
assert_eq!(resolved.branches.len(), 4);
assert!(resolved.branches.contains(&"feature/a".to_string()));
assert!(resolved.branches.contains(&"feature/b".to_string()));
assert!(resolved.branches.contains(&"feature/pay-ux".to_string()));
assert!(resolved.branches.contains(&"feature/pay-api".to_string()));
}
#[test]
fn test_cycle_detection() {
let mut parties = HashMap::new();
parties.insert("a".to_string(), Party {
members: vec!["@b".to_string()],
..Party::default()
});
parties.insert("b".to_string(), Party {
members: vec!["@c".to_string()],
..Party::default()
});
parties.insert("c".to_string(), Party {
members: vec!["@a".to_string()],
..Party::default()
});
let config = Config {
base_branch: "main".to_string(),
parties,
auto_update: crate::AutoUpdateConfig::default(),
};
let resolver = PartyResolver::new(&config);
let result = resolver.resolve_party("a");
assert!(result.is_err());
if let Err(Error::PartyCycle { .. }) = result {
} else {
panic!("Expected PartyCycle error");
}
}
#[test]
fn test_missing_party_reference() {
let config = create_test_config();
let resolver = PartyResolver::new(&config);
let result = resolver.resolve_party("nonexistent");
assert!(result.is_err());
if let Err(Error::PartyNotFound { .. }) = result {
} else {
panic!("Expected PartyNotFound error");
}
}
#[test]
fn test_dependencies() {
let config = create_test_config();
let resolver = PartyResolver::new(&config);
let deps = resolver.get_dependencies("qa").unwrap();
assert!(deps.contains(&"payments".to_string()));
assert!(!deps.contains(&"qa".to_string()));
}
}