带权有向无环图两点间最长路径求解:代码存未覆盖边缘用例
问题:求解边带权有向无环图(DAG)中源节点到汇节点的最长路径
需求:给定边带权有向无环图(DAG)、源节点、汇节点,计算源节点到汇节点的最长路径,输出路径长度及具体路径。
输入输出示例
输入
0 4 0->1:7 0->2:4 2->3:2 1->4:1 3->4:3
输出
9 0->2->3->4
现有代码问题
当前代码存在语法错误(如边拆分逻辑完全错误),且未覆盖以下边缘用例:
- 源节点与汇节点为同一节点的情况
- 源节点到汇节点不存在路径的情况
- 图中包含孤立节点的情况
修正后的代码
import sys from collections import defaultdict def topological_sort(graph, all_nodes): # 初始化入度字典,包含所有节点 indegree = {node: 0 for node in all_nodes} for node in graph: for neighbor, _ in graph[node]: indegree[neighbor] += 1 nodes_with_no_indegree = [node for node in all_nodes if indegree[node] == 0] ordering = [] while nodes_with_no_indegree: node = nodes_with_no_indegree.pop() ordering.append(node) for neighbor, _ in graph.get(node, []): indegree[neighbor] -= 1 if indegree[neighbor] == 0: nodes_with_no_indegree.append(neighbor) return ordering def longest_path(source, sink, edges): source_str = str(source) sink_str = str(sink) # 处理源汇相同的情况 if source_str == sink_str: return 0, source_str graph = defaultdict(list) all_nodes = set() # 正确拆分边信息 for edge in edges: if not edge.strip(): continue u_part, weight_part = edge.split("->") v, weight = weight_part.split(":") u = u_part.strip() v = v.strip() weight = int(weight.strip()) graph[u].append((v, -weight)) # 转为负权求最短路径 all_nodes.add(u) all_nodes.add(v) # 确保源和汇在节点集合中 all_nodes.add(source_str) all_nodes.add(sink_str) all_nodes = list(all_nodes) ordering = topological_sort(graph, all_nodes) # 初始化距离字典,源节点距离为0,其他为无穷大 dist = {node: float('inf') for node in all_nodes} dist[source_str] = 0 pred = {} for u in ordering: if dist[u] == float('inf'): continue # 源节点不可达的节点跳过 for v, weight in graph.get(u, []): if dist[v] > dist[u] + weight: dist[v] = dist[u] + weight pred[v] = u # 检查汇节点是否可达 if dist[sink_str] == float('inf'): return -1, "源节点到汇节点不存在路径" # 重建路径 path = [] current = sink_str while current in pred: path.append(current) current = pred[current] path.append(source_str) path.reverse() return -dist[sink_str], '->'.join(path) if __name__ == "__main__": source = int(sys.stdin.readline().strip()) sink = int(sys.stdin.readline().strip()) edges = [line.strip() for line in sys.stdin if line.strip()] max_val, path = longest_path(source, sink, edges) print(max_val) print(path)
代码说明
- 边拆分逻辑修正:正确按
->和:拆分边的起点、终点和权重,处理可能的空格问题 - 拓扑排序优化:传入所有节点集合,确保孤立节点也被纳入排序
- 边缘用例处理:
- 源汇相同时直接返回长度0和节点本身
- 检查汇节点是否可达,不可达时返回提示信息
- 过滤空的输入行,避免解析错误
- 路径重建优化:确保只有可达时才重建路径,避免错误拼接
内容的提问来源于stack exchange,提问作者Ali Alsawad
相关产品推荐
相关产品推荐

