[object Object]

← back to Exo

fix: Fix validation over Task types

367e76c8fad85a161588ffcf68ffbed751568e87 · 2025-07-04 17:25:14 +0100 · Arbion Halili

Files touched

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 →