← back to Exo
fix: Fix validation over Task types
367e76c8fad85a161588ffcf68ffbed751568e87 · 2025-07-04 17:25:14 +0100 · Arbion Halili
Files touched
M shared/types/events/events.pyM shared/types/states/shared.pyM shared/types/tasks/common.py
Diff
commit 367e76c8fad85a161588ffcf68ffbed751568e87
Author: Arbion Halili <99731180+ToxicPine@users.noreply.github.com>
Date: Fri Jul 4 17:25:14 2025 +0100
fix: Fix validation over Task types
---
shared/types/events/events.py | 16 ++++-----
shared/types/states/shared.py | 6 ++--
shared/types/tasks/common.py | 81 +++++++++++++++++++++++++++++++------------
3 files changed, 68 insertions(+), 35 deletions(-)
diff --git a/shared/types/events/events.py b/shared/types/events/events.py
index cd0da509..712e8936 100644
--- a/shared/types/events/events.py
+++ b/shared/types/events/events.py
@@ -1,6 +1,6 @@
from __future__ import annotations
-from typing import Any, Generic, Literal, Tuple, TypeVar
+from typing import Any, Literal, Tuple
from pydantic import BaseModel
@@ -33,9 +33,10 @@ from shared.types.profiling.common import NodePerformanceProfile
from shared.types.tasks.common import (
TaskData,
TaskId,
+ TaskState,
+ TaskStatusIncompleteType,
TaskStatusType,
TaskType,
- TaskUpdate,
)
from shared.types.worker.common import InstanceId, NodeStatus
from shared.types.worker.instances import InstanceData, InstanceStatus
@@ -54,21 +55,18 @@ class TimerData(BaseModel):
timer_id: TimerId
-TaskTypeT = TypeVar("TaskTypeT", bound=TaskType)
-
-
-class TaskCreated(Event[TaskEventTypes.TaskCreated], Generic[TaskTypeT]):
+class TaskCreated[TaskTypeT: TaskType](Event[TaskEventTypes.TaskCreated]):
event_type: Literal[TaskEventTypes.TaskCreated] = TaskEventTypes.TaskCreated
task_id: TaskId
task_data: TaskData[TaskTypeT]
- task_state: TaskUpdate[Literal[TaskStatusType.Pending], TaskTypeT]
+ task_state: TaskState[TaskTypeT, Literal[TaskStatusIncompleteType.Pending]]
on_instance: InstanceId
-class TaskUpdated(Event[TaskEventTypes.TaskUpdated], Generic[TaskTypeT]):
+class TaskUpdated[TaskTypeT: TaskType](Event[TaskEventTypes.TaskUpdated]):
event_type: Literal[TaskEventTypes.TaskUpdated] = TaskEventTypes.TaskUpdated
task_id: TaskId
- update_data: TaskUpdate[TaskStatusType, TaskTypeT]
+ update_data: TaskState[TaskTypeT, TaskStatusType]
class TaskDeleted(Event[TaskEventTypes.TaskDeleted]):
diff --git a/shared/types/states/shared.py b/shared/types/states/shared.py
index e366602f..15caa2d0 100644
--- a/shared/types/states/shared.py
+++ b/shared/types/states/shared.py
@@ -5,7 +5,7 @@ from pydantic import BaseModel
from shared.types.common import NodeId
from shared.types.events.common import InstanceStateEventTypes, State, TaskEventTypes
-from shared.types.tasks.common import Task, TaskId, TaskType
+from shared.types.tasks.common import Task, TaskId, TaskStatusType, TaskType
from shared.types.worker.common import InstanceId
from shared.types.worker.instances import BaseInstance
@@ -15,7 +15,7 @@ class KnownInstances(State[InstanceStateEventTypes]):
class Tasks(State[TaskEventTypes]):
- tasks: Mapping[TaskId, Task[TaskType]]
+ tasks: Mapping[TaskId, Task[TaskType, TaskStatusType]]
class SharedState(BaseModel):
@@ -27,4 +27,4 @@ class SharedState(BaseModel):
def get_tasks_by_instance(
self, instance_id: InstanceId
- ) -> Sequence[Task[TaskType]]: ...
+ ) -> Sequence[Task[TaskType, TaskStatusType]]: ...
diff --git a/shared/types/tasks/common.py b/shared/types/tasks/common.py
index a01c641d..114c0550 100644
--- a/shared/types/tasks/common.py
+++ b/shared/types/tasks/common.py
@@ -1,9 +1,9 @@
from collections.abc import Mapping
from enum import Enum
-from typing import Generic, Literal, TypeVar, Union
+from typing import Annotated, Generic, Literal, TypeVar, Union
import openai.types.chat as openai
-from pydantic import BaseModel
+from pydantic import BaseModel, Field, TypeAdapter
from shared.types.common import NewUUID
from shared.types.worker.common import InstanceId, RunnerId
@@ -18,11 +18,10 @@ class TaskType(str, Enum):
ChatCompletionStreaming = "ChatCompletionStreaming"
-TaskTypeT = TypeVar("TaskTypeT", bound=TaskType)
+TaskTypeT = TypeVar("TaskTypeT", bound=TaskType, covariant=True)
-class TaskData(BaseModel, Generic[TaskTypeT]):
- task_type: TaskTypeT
+class TaskData(BaseModel, Generic[TaskTypeT]): ...
class ChatCompletionNonStreamingTask(TaskData[TaskType.ChatCompletionNonStreaming]):
@@ -39,48 +38,84 @@ class ChatCompletionStreamingTask(TaskData[TaskType.ChatCompletionStreaming]):
task_data: openai.completion_create_params.CompletionCreateParams
-class TaskStatusType(str, Enum):
+class TaskStatusIncompleteType(str, Enum):
Pending = "Pending"
Running = "Running"
Failed = "Failed"
+
+
+class TaskStatusCompleteType(str, Enum):
Complete = "Complete"
-TaskStatusTypeT = TypeVar(
- "TaskStatusTypeT", bound=Union[TaskStatusType, Literal["Complete"]]
-)
+TaskStatusType = Union[TaskStatusIncompleteType, TaskStatusCompleteType]
+
+TaskStatusTypeT = TypeVar("TaskStatusTypeT", bound=TaskStatusType, covariant=True)
-class TaskArtifact(BaseModel, Generic[TaskTypeT]): ...
+class TaskArtifact[TaskTypeT: TaskType, TaskStatusTypeT: TaskStatusType](BaseModel): ...
-class TaskUpdate(BaseModel, Generic[TaskStatusTypeT, TaskTypeT]):
+
+class IncompleteTaskArtifact[TaskTypeT: TaskType](
+ TaskArtifact[TaskTypeT, TaskStatusIncompleteType]
+):
+ pass
+
+
+class TaskStatusUpdate[TaskStatusTypeT: TaskStatusType](BaseModel):
task_status: TaskStatusTypeT
-class PendingTask(TaskUpdate[TaskStatusType.Pending, TaskTypeT]):
- task_status: Literal[TaskStatusType.Pending]
+class PendingTaskStatus(TaskStatusUpdate[TaskStatusIncompleteType.Pending]):
+ task_status: Literal[TaskStatusIncompleteType.Pending] = (
+ TaskStatusIncompleteType.Pending
+ )
-class RunningTask(TaskUpdate[TaskStatusType.Running, TaskTypeT]):
- task_status: Literal[TaskStatusType.Running]
+class RunningTaskStatus(TaskStatusUpdate[TaskStatusIncompleteType.Running]):
+ task_status: Literal[TaskStatusIncompleteType.Running] = (
+ TaskStatusIncompleteType.Running
+ )
-class CompletedTask(TaskUpdate[TaskStatusType.Complete, TaskTypeT]):
- task_status: Literal[TaskStatusType.Complete]
- task_artifact: TaskArtifact[TaskTypeT]
+class CompletedTaskStatus(TaskStatusUpdate[TaskStatusCompleteType.Complete]):
+ task_status: Literal[TaskStatusCompleteType.Complete] = (
+ TaskStatusCompleteType.Complete
+ )
-class FailedTask(TaskUpdate[TaskStatusType.Failed, TaskTypeT]):
- task_status: Literal[TaskStatusType.Failed]
+class FailedTaskStatus(TaskStatusUpdate[TaskStatusIncompleteType.Failed]):
+ task_status: Literal[TaskStatusIncompleteType.Failed] = (
+ TaskStatusIncompleteType.Failed
+ )
error_message: Mapping[RunnerId, str]
-class BaseTask(BaseModel, Generic[TaskTypeT]):
+class TaskState(BaseModel, Generic[TaskTypeT, TaskStatusTypeT]):
+ task_status: TaskStatusUpdate[TaskStatusTypeT]
+ task_artifact: TaskArtifact[TaskTypeT, TaskStatusTypeT]
+
+
+class BaseTask(BaseModel, Generic[TaskTypeT, TaskStatusTypeT]):
+ task_type: TaskTypeT
task_data: TaskData[TaskTypeT]
- task_status: TaskUpdate[TaskStatusType, TaskTypeT]
+ task_state: TaskState[TaskTypeT, TaskStatusTypeT]
on_instance: InstanceId
-class Task(BaseTask[TaskTypeT]):
+BaseTaskAnnotated = Annotated[
+ Union[
+ BaseTask[Literal[TaskType.ChatCompletionNonStreaming], TaskStatusType],
+ BaseTask[Literal[TaskType.ChatCompletionStreaming], TaskStatusType],
+ ],
+ Field(discriminator="task_type"),
+]
+
+BaseTaskValidator: TypeAdapter[BaseTask[TaskType, TaskStatusType]] = TypeAdapter(
+ BaseTaskAnnotated
+)
+
+
+class Task(BaseTask[TaskTypeT, TaskStatusTypeT]):
task_id: TaskId
← cda3de2a fix: Use state for tasks
·
back to Exo
·
Matt's interfaces 03a1cf59 →