如何解决将PipelineTask实例传入自身add方法时报NameError的问题
问题描述
我想要将PipelineTask类的实例传入该类自身的PipelineTask.add方法,但运行时触发NameError,提示PipelineTask未定义。
我推测该问题是因为PipelineTask要在PipelineTask.__init__()调用完成后才会完成绑定。
class Task(BaseModel, abc.ABC): id: str @abc.abstractmethod async def run(self): pass class PipelineTask(Task): @abc.abstractmethod async def run(self): pass def add(self, task: PipelineTask): ## TODO: how do a pass a instance of self here? next = self if self.next == None: self.next = task next = self.next else: current = self.next while current != None: current = current.next next = current current.next = next return self class Pipeline(BaseModel): """ Pipeline execute a sequence of tasks ... init = PipelineTask(id=0) pipeline = Pipeline(init=PrepareDataPipelineTask(id='prepare')) pipeline.add(ExecutePipelineTask(id='execute')).add(CollectResultsPipelineTask(id='collect')) pipeline.run() ... """ # The pipelines innitial task init: PipelineTask async def run(self): await self.init.run() has_next = self.init.next != None next = self.init while has_next: next = next.next await next.run() has_next = next.next != None ## Adds a task to the end of the pipeline async def add(self, next: PipelineTask): """add Task to the end of the pipeline""" self.init.add(next) class StdoutTask(PipelineTask): async def run(self): print(f"[Worker {self.id}] testing...") async def test_create_pipeline(): tasks = ( StdoutTask(id=1, next=None) .add(StdoutTask(id=2, next=None)) .add(StdoutTask(id=3, next=None)) ) pipeline = Pipeline(init=tasks) await pipeline.run()
示例用法:
class StdoutTask(PipelineTask): async def run(self): print(f"[Worker {self.id}] testing...") @pytest.mark.asyncio async def test_create_pipeline(): tasks = StdoutTask(id=1).add(StdoutTask(id=2)).add(StdoutTask(id=3)) pipeline = Pipeline(init=tasks) await pipeline.run() pass
我曾尝试过两种解决方法都没有成功:
- 移除
task的类型标注,运行后触发AttributeError,提示对象不存在next属性:
def add(self, task): ...
- 修改
task.__class__ = PipelineTask,但该操作只能新增对应方法,无法补全属性。
以下是可单文件运行的问题复现代码:
from pydantic import BaseModel import abc import asyncio class Task(BaseModel, abc.ABC): id: str @abc.abstractmethod async def run(self): pass class PipelineTask(Task): @abc.abstractmethod async def run(self): pass def add(self, task: PipelineTask): ## TODO: how do a pass a instance of self here? next = self if self.next == None: self.next = task next = self.next else: current = self.next while current != None: current = current.next next = current current.next = next return self class Pipeline(BaseModel): """ Pipeline execute a sequence of tasks ... init = PipelineTask(id=0) pipeline = Pipeline(init=PrepareDataPipelineTask(id='prepare')) pipeline.add(ExecutePipelineTask(id='execute')).add(CollectResultsPipelineTask(id='collect')) pipeline.run() ... """ # The pipelines innitial task init: PipelineTask async def run(self): await self.init.run() has_next = self.init.next != None next = self.init while has_next: next = next.next await next.run() has_next = next.next != None ## Adds a task to the end of the pipeline async def add(self, next: PipelineTask): """add Task to the end of the pipeline""" self.init.add(next) class StdoutTask(PipelineTask): async def run(self): print(f"[Worker {self.id}] testing...") async def test_create_pipeline(): tasks = ( StdoutTask(id=1, next=None) .add(StdoutTask(id=2, next=None)) .add(StdoutTask(id=3, next=None)) ) pipeline = Pipeline(init=tasks) await pipeline.run()
解决方案
结合字符串形式的类型前向引用解决NameError,同时使用getattr读取属性避免AttributeError,修改后的代码如下:
def add(self, task: "PipelineTask"): next = getattr(self, "next", None) if self.next == None: self.next = task next = self.next else: current = getattr(self, "next", None) while current != None: current = getattr(current, "next", None) next = getattr(current, "next", None) current.next = next return self
内容的提问来源于stack exchange,提问作者Babbleshack
相关产品推荐
相关产品推荐

