You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python双向链表单元测试异常求助:迭代器与长度计算错误

双向链表实现问题排查

问题描述

我需要实现一个包含以下4种操作的双向链表:

  • push:尾部插入值;
  • pop:尾部删除值;
  • shift:头部删除值;
  • unshift:头部插入值;

目前代码已实现上述操作,但单元测试出现两个问题:

  1. shift方法执行后,链表长度未正确更新;
  2. 链表对象无法被迭代。

现有代码实现

双向链表代码

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.18 18:45:47