如何用Python递归实现Haskell的惰性递归版span函数
Haskell span 函数的Python惰性递归实现
span函数的作用是将输入序列拆分为两部分:第一部分是满足给定谓词的最长前缀,第二部分是序列中剩余的所有元素,语义和Haskell标准库的
span完全对齐。
实现代码
def span(p, iterable): iterator = iter(iterable) mismatch_elem = None def matched_gen(): nonlocal mismatch_elem try: elem = next(iterator) except StopIteration: return if p(elem): yield elem yield from matched_gen() else: mismatch_elem = elem return def remaining_gen(): nonlocal mismatch_elem if mismatch_elem is not None: yield mismatch_elem yield from iterator return (matched_gen(), remaining_gen())
实现特性
- 完全惰性:仅在迭代返回的生成器时才消费输入元素,支持处理无限长可迭代对象
- 递归逻辑:匹配前缀的生成器采用递归实现,符合要求的风格,
yield from简化了递归生成器的元素传递 - 无额外缓存:利用Python迭代器的内置状态特性,不需要缓存整个输入序列,内存效率更高
- 语义完全对齐Haskell实现:
- 空输入返回两个空生成器
- 第一个元素不满足谓词时,第一个生成器为空,第二个生成器返回完整输入序列
使用示例
# 测试用例1:普通列表 is_less_than_3 = lambda x: x < 3 test_list = [1, 2, 3, 4, 1, 2] matched, remaining = span(is_less_than_3, test_list) print(list(matched)) # 输出: [1, 2] print(list(remaining)) # 输出: [3, 4, 1, 2] # 测试用例2:首个元素不匹配 is_greater_than_5 = lambda x: x > 5 matched, remaining = span(is_greater_than_5, [1,2,3]) print(list(matched)) # 输出: [] print(list(remaining)) # 输出: [1,2,3] # 测试用例3:空输入 matched, remaining = span(is_less_than_3, []) print(list(matched)) # 输出: [] print(list(remaining)) # 输出: []
内容的提问来源于stack exchange,提问作者Eric Auld
相关产品推荐
相关产品推荐

