Python双向链表单元测试异常求助:迭代器与长度计算错误
双向链表实现问题排查
问题描述
我需要实现一个包含以下4种操作的双向链表:
push:尾部插入值;pop:尾部删除值;shift:头部删除值;unshift:头部插入值;
目前代码已实现上述操作,但单元测试出现两个问题:
shift方法执行后,链表长度未正确更新;- 链表对象无法被迭代。
现有代码实现
双向链表代码
class Node(object): def __init__(self, value, succeeding=None, previous=None): # reference to next node into the doubly linked list self.next = succeeding # reference to previous node into the doubly linked list self.prev = previous # adding an value element to the node self.data = value class LinkedList(object): def __init__(self): self.head = None self.tail = None def push(self, new_data): """ insert value at back in the list Args: new_data (object): the value to insert at the back in the list Returns: None """ # creating a new node with the desired value new_node = Node(new_data) # newly created node's next pointer will refer to the old head new_node.next = self.head # Checks whether list is empty or not if self.head != None: # old head's previous pointer will refer to newly created node self.head.prev = new_node # new node becomes the new head self.head = new_node new_node.prev = None else: # If the list is empty, make new node both head and tail self.head = new_node self.tail = new_node new_node.prev = None # There's only one element so both pointers refer to null def unshift(self, new_data): """ Insert value at front in the list Args: new_data (object): the value to insert at the front in the list Returns: None """ new_node = Node(new_data) new_node.prev = self.tail # checks whether the list is empty, if so make both head and tail as new node if self.tail == None: self.head = new_node self.tail = new_node # the first element's previous pointer has to refer to null new_node.next = None # If list is not empty, change pointers accordingly else: self.tail.next = new_node new_node.next = None self.tail = new_node # Make new node the new tail def pop(self): """ Remove value at back in the list Args: None Returns: data (object): The value removed at the back in the list """ if self.head == None: print("List is empty") else: temp = self.head # temp.next.prev = None # remove previous pointer referring to old head self.head = temp.next # make second element the new head temp.next = None # remove next pointer referring to new head return temp.data def shift(self): """ Remove value at front in the list Args: None Returns: data (object): The value removed at the front in the list """ if self.tail == None: print("List is empty") else: temp = self.tail # temp.prev.next = None # removes next pointer referring to old tail self.tail = temp.prev # make second to last element the new tail temp.prev = None # remove previous pointer referring to new tail return temp.data def taille(self, dll): """ Get the length of the doubly linked list Args: dll (Object): the doubly linked list Returns: length (list): The length of the DLL """ n = 0 lst = dll while lst is not None: n = n + 1 lst = lst.next return n def __len__(self): return self.taille(self.head) # This function prints contents of linked list # starting from the given node def printList(self, node): print("\nTraversal in forward direction") while node: print(" {}".format(node.data), end="") last = node node = node.next # Start with empty list if __name__ == "__main__": llist = LinkedList() # Insert 6. So the list becomes 6->None llist.unshift(6) # Insert 7 at the beginning. # So linked list becomes 7->6->None llist.push(7) # Insert 1 at the beginning. # So linked list becomes 1->7->6->None llist.push(1) # Insert 4 at the end. # So linked list becomes 1->7->6->4->None llist.unshift(4) print("Created DLL is: ") llist.printList(llist.head) print('\n', len(llist)) print('shift value return ',llist.shift()) print('pop value return ',llist.pop()) print("\n List after pop/shift element: ") llist.printList(llist.head) # print('\n', len(llist))
单元测试用例
import unittest from linked_list import LinkedList class LinkedListTest(unittest.TestCase): def test_push_pop(self): lst = LinkedList() lst.push(10) lst.push(20) self.assertEqual(lst.pop(), 20) self.assertEqual(lst.pop(), 10) def test_push_shift(self): lst = LinkedList() lst.push(10) lst.push(20) self.assertEqual(lst.shift(), 10) self.assertEqual(lst.shift(), 20) def test_unshift_shift(self): lst = LinkedList() lst.unshift(10) lst.unshift(20) self.assertEqual(lst.shift(), 20) self.assertEqual(lst.shift(), 10) def test_unshift_pop(self): lst = LinkedList() lst.unshift(10) lst.unshift(20) self.assertEqual(lst.pop(), 10) self.assertEqual(lst.pop(), 20) def test_all(self): lst = LinkedList() lst.push(10) lst.push(20) self.assertEqual(lst.pop(), 20) lst.push(30) self.assertEqual(lst.shift(), 10) lst.unshift(40) lst.push(50) self.assertEqual(lst.shift(), 40) self.assertEqual(lst.pop(), 50) self.assertEqual(lst.shift(), 30) def test_length(self): lst = LinkedList() lst.push(10) lst.push(20) self.assertEqual(len(lst), 2) lst.shift() self.assertEqual(len(lst), 1) lst.pop() self.assertEqual(len(lst), 0) def test_iterator(self): lst = LinkedList() lst.push(10) lst.push(20) iterator = iter(lst) self.assertEqual(next(iterator), 10) self.assertEqual(next(iterator), 20) def supprime(self): self.supprimeCellule(self.head, 1) def supprimeCellule(self, L, n): """ Supprime la cellule après la nième cellule dans la liste L""" if n > self.taille(L)-1: raise IndexError("Indice invalide") for i in range(n-1): L = L.next L.next = L.next.next if __name__ == '__main__': unittest.main()
错误信息
.EF.... ====================================================================== ERROR: test_iterator (__main__.LinkedListTest) ---------------------------------------------------------------------- Traceback (most recent call last): File "C:\Users\guera\OneDrive\Documents\Python Scripts\preludd-recrutement-overlap\linked-list\linked_list_test.py", line 62, in test_iterator iterator = iter(lst) TypeError: 'LinkedList' object is not iterable ====================================================================== FAIL: test_length (__main__.LinkedListTest) ---------------------------------------------------------------------- Traceback (most recent call last): File "C:\Users\guera\OneDrive\Documents\Python Scripts\preludd-recrutement-overlap\linked-list\linked_list_test.py", line 54, in test_length self.assertEqual(len(lst), 1) AssertionError: 2 != 1 ---------------------------------------------------------------------- Ran 7 tests in 0.002s FAILED (failures=1, errors=1)
问题修复方案
1. 修正方法功能与实现的颠倒问题
原代码中push/unshift、pop/shift的功能和实现完全反向,需重新实现各方法:
修正push(尾部插入)
def push(self, new_data): """ insert value at back in the list Args: new_data (object): the value to insert at the back in the list Returns: None """ new_node = Node(new_data) new_node.next = None # 新节点是尾部,next指向None if self.tail is not None: self.tail.next = new_node new_node.prev = self.tail self.tail = new_node else: # 空链表时,head和tail都指向新节点 self.head = new_node self.tail = new_node new_node.prev = None
修正unshift(头部插入)
def unshift(self, new_data): """ Insert value at front in the list Args: new_data (object): the value to insert at the front in the list Returns: None """ new_node = Node(new_data) new_node.prev = None # 新节点是头部,prev指向None if self.head is not None: self.head.prev = new_node new_node.next = self.head self.head = new_node else: # 空链表时,head和tail都指向新节点 self.head = new_node self.tail = new_node new_node.next = None
修正pop(尾部删除)
def pop(self): """ Remove value at back in the list Args: None Returns: data (object): The value removed at the back in the list """ if self.tail is None: print("List is empty") return None temp = self.tail if self.head == self.tail: # 只剩一个节点时,清空head和tail self.head = None self.tail = None else: self.tail = temp.prev self.tail.next = None # 断开新尾部与旧尾部的连接 temp.prev = None return temp.data
修正shift(头部删除)
def shift(self): """ Remove value at front in the list Args: None Returns: data (object): The value removed at the front in the list """ if self.head is None: print("List is empty") return None temp = self.head if self.head == self.tail: # 只剩一个节点时,清空head和tail self.head = None self.tail = None else: self.head = temp.next self.head.prev = None # 断开新头部与旧头部的连接 temp.next = None return temp.data
2. 实现迭代器协议
添加__iter__方法,让链表支持迭代:
def __iter__(self): current = self.head while current is not None: yield current.data current = current.next
3. 优化长度计算(可选)
原taille方法需要遍历整个链表计算长度,可在类中维护一个length属性,插入/删除时更新,提升效率:
class LinkedList(object): def __init__(self): self.head = None self.tail = None self.length = 0 # 添加长度属性 # 在push方法末尾添加: self.length += 1 # 在unshift方法末尾添加: self.length += 1 # 在pop方法中,删除节点后添加: self.length -= 1 # 在shift方法中,删除节点后添加: self.length -= 1 # 修改__len__方法: def __len__(self): return self.length
修复后测试结果
所有单元测试均可通过,包括长度校验和迭代功能。
内容的提问来源于stack exchange,提问作者Serge
相关产品推荐
相关产品推荐

