diff --git a/CLAUDE.md b/CLAUDE.md index 3d1522d2..c4467d84 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -380,7 +380,7 @@ Use the escape hatch for "this needs a maintenance window with traffic shifted a ##### Anti-pattern → linter rule reference -The migration linter at `agentex/scripts/lint_migrations.py` enforces these rules at PR time via `.github/workflows/migration-lint.yml`. It only checks files changed vs the PR base, so existing migrations are not retro-flagged. The mapping below is what the linter catches: +The migration linter at `agentex/scripts/ci_tools/migration_lint.py` enforces these rules at PR time via `.github/workflows/migration-lint.yml`. It only checks files changed vs the PR base, so existing migrations are not retro-flagged. The mapping below is what the linter catches: | Anti-pattern | Linter rule | |---|---| @@ -393,7 +393,7 @@ The migration linter at `agentex/scripts/lint_migrations.py` enforces these rule Run the linter locally before pushing: ```bash -agentex/scripts/lint_migrations.py --base-ref origin/main +agentex/scripts/ci_tools/migration_lint.py --base origin/main ``` ##### Other rules diff --git a/agentex/database/migrations/alembic/versions/2026_07_22_1200_add_task_current_state_b2c3d4e5f6a7.py b/agentex/database/migrations/alembic/versions/2026_07_22_1200_add_task_current_state_b2c3d4e5f6a7.py new file mode 100644 index 00000000..468e51d5 --- /dev/null +++ b/agentex/database/migrations/alembic/versions/2026_07_22_1200_add_task_current_state_b2c3d4e5f6a7.py @@ -0,0 +1,28 @@ +"""add task current_state + +Revision ID: b2c3d4e5f6a7 +Revises: a1b2c3d4e5f6 +Create Date: 2026-07-22 12:00:00.000000 + +""" +from typing import Sequence, Union + +from alembic import op + + +# revision identifiers, used by Alembic. +revision: str = 'b2c3d4e5f6a7' +down_revision: Union[str, None] = 'a1b2c3d4e5f6' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # Nullable additive column; idempotent, metadata-only, non-blocking. + op.execute( + "ALTER TABLE tasks ADD COLUMN IF NOT EXISTS current_state VARCHAR(255)" + ) + + +def downgrade() -> None: + op.execute("ALTER TABLE tasks DROP COLUMN IF EXISTS current_state") diff --git a/agentex/openapi.yaml b/agentex/openapi.yaml index b62ae7a1..5aa6258c 100644 --- a/agentex/openapi.yaml +++ b/agentex/openapi.yaml @@ -6622,18 +6622,24 @@ components: - type: 'null' title: The timestamp when the task's content was cleaned for retention compliance; null when active - params: + task_metadata: anyOf: - additionalProperties: true type: object - type: 'null' - title: Task parameters - task_metadata: + title: Task metadata + current_state: + anyOf: + - type: string + - type: 'null' + title: Opaque label mirroring the agent's StateMachine current state; null + when the agent does not emit one. Orthogonal to 'status'. + params: anyOf: - additionalProperties: true type: object - type: 'null' - title: Task metadata + title: Task parameters type: object required: - id @@ -6845,18 +6851,24 @@ components: - type: 'null' title: The timestamp when the task's content was cleaned for retention compliance; null when active - params: + task_metadata: anyOf: - additionalProperties: true type: object - type: 'null' - title: Task parameters - task_metadata: + title: Task metadata + current_state: + anyOf: + - type: string + - type: 'null' + title: Opaque label mirroring the agent's StateMachine current state; null + when the agent does not emit one. Orthogonal to 'status'. + params: anyOf: - additionalProperties: true type: object - type: 'null' - title: Task metadata + title: Task parameters agents: anyOf: - items: @@ -6936,6 +6948,12 @@ components: type: object - type: 'null' title: Task metadata + current_state: + anyOf: + - type: string + - type: 'null' + title: Opaque label mirroring the agent's StateMachine current state; null + when the agent does not emit one. Orthogonal to 'status'. agents: anyOf: - items: @@ -7504,6 +7522,12 @@ components: - type: 'null' title: Optional shallow-merge patch applied to the task's params column. Top-level keys overwrite; pass full nested objects to change subfields. + current_state: + anyOf: + - type: string + maxLength: 255 + - type: 'null' + title: If provided, replaces the task's current_state label. type: object title: UpdateTaskRequest ValidationError: diff --git a/agentex/src/adapters/orm.py b/agentex/src/adapters/orm.py index e5f7b139..0851afad 100644 --- a/agentex/src/adapters/orm.py +++ b/agentex/src/adapters/orm.py @@ -24,6 +24,7 @@ from src.domain.entities.deployments import DeploymentStatus from src.domain.entities.tasks import TaskStatus from src.utils.ids import orm_id +from src.utils.task_constants import CURRENT_STATE_MAX_LENGTH BaseORM = declarative_base() @@ -75,6 +76,8 @@ class TaskORM(BaseORM): cleaned_at = Column(DateTime(timezone=True), nullable=True) params = Column(JSONB, nullable=True) task_metadata = Column(JSONB, nullable=True) + # Opaque agent-state label, orthogonal to `status`; capped since it rides every task_updated SSE payload. + current_state = Column(String(CURRENT_STATE_MAX_LENGTH), nullable=True) # Many-to-Many relationship with agents agents = relationship("AgentORM", secondary="task_agents", back_populates="tasks") diff --git a/agentex/src/api/routes/tasks.py b/agentex/src/api/routes/tasks.py index 6d3607c7..79561998 100644 --- a/agentex/src/api/routes/tasks.py +++ b/agentex/src/api/routes/tasks.py @@ -200,6 +200,7 @@ async def update_task( id=task_id, task_metadata=request.task_metadata, merge_params=request.merge_params, + current_state=request.current_state, ) return Task.model_validate(updated_task_entity) @@ -221,6 +222,7 @@ async def update_task_by_name( name=task_name, task_metadata=request.task_metadata, merge_params=request.merge_params, + current_state=request.current_state, ) return Task.model_validate(updated_task_entity) diff --git a/agentex/src/api/schemas/tasks.py b/agentex/src/api/schemas/tasks.py index f5efaad6..97dce714 100644 --- a/agentex/src/api/schemas/tasks.py +++ b/agentex/src/api/schemas/tasks.py @@ -6,6 +6,7 @@ from src.api.schemas.agents import Agent from src.utils.model_utils import BaseModel +from src.utils.task_constants import CURRENT_STATE_MAX_LENGTH class TaskRelationships(str, Enum): @@ -26,45 +27,34 @@ class TaskStatus(str, Enum): DELETED = "DELETED" -class Task(BaseModel): - id: str = Field( - ..., - title="Unique Task ID", - ) - name: str | None = Field( - None, - title="Unique name of the task", - ) - status: TaskStatus | None = Field( - None, - title="The current status of the task", - ) - status_reason: str | None = Field( - None, - title="The reason for the current task status", - ) - created_at: datetime | None = Field( - None, - title="The timestamp when the task was created", - ) - updated_at: datetime | None = Field( - None, - title="The timestamp when the task was last updated", - ) +class _TaskBase(BaseModel): + """Shared fields for Task and TaskSummary (everything except `params`).""" + + id: str = Field(..., title="Unique Task ID") + name: str | None = Field(None, title="Unique name of the task") + status: TaskStatus | None = Field(None, title="The current status of the task") + status_reason: str | None = Field(None, title="The reason for the current task status") + created_at: datetime | None = Field(None, title="The timestamp when the task was created") + updated_at: datetime | None = Field(None, title="The timestamp when the task was last updated") cleaned_at: datetime | None = Field( None, title="The timestamp when the task's content was cleaned for retention compliance; null when active", ) - params: dict[str, Any] | None = Field( - None, - title="Task parameters", - ) - task_metadata: dict[str, Any] | None = Field( + task_metadata: dict[str, Any] | None = Field(None, title="Task metadata") + # Writes are bounded; reads are not, so widening the column won't 500. + current_state: str | None = Field( None, - title="Task metadata", + title=( + "Opaque label mirroring the agent's StateMachine current state; " + "null when the agent does not emit one. Orthogonal to 'status'." + ), ) +class Task(_TaskBase): + params: dict[str, Any] | None = Field(None, title="Task parameters") + + class TaskResponse(Task): """Task response model with optional related data based on relationships""" @@ -74,43 +64,11 @@ class TaskResponse(Task): ) -class TaskSummary(BaseModel): +class TaskSummary(_TaskBase): """Lean list-response shape. Omits `params` (the arbitrary create-time payload, which can carry per-caller secrets and PII); fetch GET /tasks/{id} for the full record.""" - id: str = Field( - ..., - title="Unique Task ID", - ) - name: str | None = Field( - None, - title="Unique name of the task", - ) - status: TaskStatus | None = Field( - None, - title="The current status of the task", - ) - status_reason: str | None = Field( - None, - title="The reason for the current task status", - ) - created_at: datetime | None = Field( - None, - title="The timestamp when the task was created", - ) - updated_at: datetime | None = Field( - None, - title="The timestamp when the task was last updated", - ) - cleaned_at: datetime | None = Field( - None, - title="The timestamp when the task's content was cleaned for retention compliance; null when active", - ) - task_metadata: dict[str, Any] | None = Field( - None, - title="Task metadata", - ) agents: list["Agent"] | None = Field( default=None, title="Agents associated with this task (only populated when 'agents' view is requested)", @@ -130,6 +88,11 @@ class UpdateTaskRequest(BaseModel): "subfields." ), ) + current_state: str | None = Field( + None, + max_length=CURRENT_STATE_MAX_LENGTH, + title="If provided, replaces the task's current_state label.", + ) class TaskStatusReasonRequest(BaseModel): diff --git a/agentex/src/domain/entities/tasks.py b/agentex/src/domain/entities/tasks.py index 21949a7d..cd6c798f 100644 --- a/agentex/src/domain/entities/tasks.py +++ b/agentex/src/domain/entities/tasks.py @@ -73,6 +73,13 @@ class TaskEntity(BaseModel): None, title="Task metadata", ) + current_state: str | None = Field( + None, + title=( + "Opaque label mirroring the agent's StateMachine current state; " + "null when the agent does not emit one. Orthogonal to 'status'." + ), + ) # allow extra fields for agents relationships model_config = ConfigDict(extra="allow") @@ -80,15 +87,4 @@ class TaskEntity(BaseModel): def convert_task_to_entity(task: Task) -> TaskEntity: """Converts the pydantic model from the API layer to the domain layer""" - - return TaskEntity( - id=task.id, - name=task.name, - status=TaskStatus[task.status.value] if task.status is not None else None, - status_reason=task.status_reason, - created_at=task.created_at, - updated_at=task.updated_at, - cleaned_at=task.cleaned_at, - params=task.params, - task_metadata=task.task_metadata, - ) + return TaskEntity.model_validate(task) diff --git a/agentex/src/domain/repositories/task_repository.py b/agentex/src/domain/repositories/task_repository.py index 45662015..eb7e06e3 100644 --- a/agentex/src/domain/repositories/task_repository.py +++ b/agentex/src/domain/repositories/task_repository.py @@ -1,6 +1,6 @@ from collections.abc import Sequence from datetime import UTC, datetime, timedelta -from typing import Annotated, Literal +from typing import Annotated, Any, Literal from fastapi import Depends from sqlalchemy import cast, distinct, func, select, update @@ -21,6 +21,9 @@ logger = make_logger(__name__) +# Columns update_mutable_fields is allowed to set (status/params have their own atomic paths). +_MUTABLE_TASK_COLUMNS = frozenset({"task_metadata", "current_state"}) + class TaskRepository(PostgresCRUDRepository[TaskORM, TaskEntity, TaskRelationships]): """Repository for Task entity with relationship loading support""" @@ -231,20 +234,36 @@ async def merge_params(self, task_id: str, patch: dict) -> TaskEntity | None: general-purpose updater. """ + # ``COALESCE(params, '{}'::jsonb)`` so a NULL existing value doesn't poison the + # concat; explicit JSONB casts so Postgres picks the jsonb ``||`` (not text concat). + existing = func.coalesce(TaskORM.params, cast({}, JSONB)) + merged = existing.op("||", return_type=JSONB)(cast(patch, JSONB)) + return await self._update_returning(task_id, {"params": merged}) + + async def update_mutable_fields( + self, task_id: str, fields: dict[str, Any] + ) -> TaskEntity | None: + """Column-scoped atomic update; can't clobber status/params. Returns updated entity or None.""" + unknown = fields.keys() - _MUTABLE_TASK_COLUMNS + if unknown: + raise ValueError( + f"update_mutable_fields may only set {sorted(_MUTABLE_TASK_COLUMNS)}; " + f"got disallowed columns {sorted(unknown)}" + ) + return await self._update_returning(task_id, fields) + + async def _update_returning( + self, task_id: str, values: dict[str, Any] + ) -> TaskEntity | None: + """UPDATE … SET … WHERE id → updated entity or None. Shared by merge_params and update_mutable_fields.""" async with ( self.start_async_db_session(True) as session, async_sql_exception_handler(), ): - # ``COALESCE(params, '{}'::jsonb)`` so a NULL existing value - # doesn't poison the concat to NULL. Both operands cast to - # JSONB explicitly so Postgres picks the JSONB ``||`` operator - # (not the text concat overload). - existing = func.coalesce(TaskORM.params, cast({}, JSONB)) - merged = existing.op("||", return_type=JSONB)(cast(patch, JSONB)) stmt = ( update(TaskORM) .where(TaskORM.id == task_id) - .values(params=merged) + .values(**values) .returning(TaskORM) ) result = await session.execute(stmt) diff --git a/agentex/src/domain/services/task_service.py b/agentex/src/domain/services/task_service.py index 26e8437f..300abb5a 100644 --- a/agentex/src/domain/services/task_service.py +++ b/agentex/src/domain/services/task_service.py @@ -205,19 +205,7 @@ async def transition_task_status( if updated_task is None: return None - try: - topic = get_task_event_stream_topic(task_id=task_id) - await self.stream_repository.send_data( - topic, - TaskStreamTaskUpdatedEventEntity( - type="task_updated", task=updated_task - ).model_dump(mode="json"), - ) - logger.info(f"task_updated event published to topic: {topic}") - except Exception as e: - logger.error( - f"Error sending task_updated event to stream: {e}", exc_info=True - ) + await self._publish_task_updated(updated_task) return updated_task @@ -227,13 +215,29 @@ async def update_task(self, task: TaskEntity) -> TaskEntity: """ updated_task = await self.task_repository.update(task) + await self._publish_task_updated(updated_task) + + return updated_task + + async def update_mutable_fields( + self, task_id: str, fields: dict[str, Any] + ) -> TaskEntity | None: + """Column-scoped atomic update, then publish task_updated. Returns updated entity or None.""" + updated_task = await self.task_repository.update_mutable_fields(task_id, fields) + if updated_task is None: + return None + + await self._publish_task_updated(updated_task) + + return updated_task + + async def _publish_task_updated(self, task: TaskEntity) -> None: try: - # The Redis adapter now handles binary data properly topic = get_task_event_stream_topic(task_id=task.id) await self.stream_repository.send_data( topic, TaskStreamTaskUpdatedEventEntity( - type="task_updated", task=updated_task + type="task_updated", task=task ).model_dump(mode="json"), ) logger.info(f"task_updated event published to topic: {topic}") @@ -242,8 +246,6 @@ async def update_task(self, task: TaskEntity) -> TaskEntity: f"Error sending task_updated event to stream: {e}", exc_info=True ) - return updated_task - async def merge_task_params(self, task_id: str, patch: dict) -> TaskEntity | None: """Atomically shallow-merge ``patch`` into ``tasks.params``. Returns the updated entity, or ``None`` if no task with ``task_id`` exists. diff --git a/agentex/src/domain/use_cases/tasks_use_case.py b/agentex/src/domain/use_cases/tasks_use_case.py index 8cc619a1..08de16d6 100644 --- a/agentex/src/domain/use_cases/tasks_use_case.py +++ b/agentex/src/domain/use_cases/tasks_use_case.py @@ -100,8 +100,9 @@ async def update_mutable_fields_on_task( name: str | None = None, task_metadata: dict[str, Any] | None = None, merge_params: dict[str, Any] | None = None, + current_state: str | None = None, ) -> TaskEntity: - """Update mutable fields on a task entity. This is used by our API since not all fields should be mutable.""" + """Update mutable fields on a task; ``None`` for any field means "not supplied".""" if not id and not name: raise ClientError("Either id or name must be provided") @@ -114,16 +115,9 @@ async def update_mutable_fields_on_task( else: raise ItemDoesNotExist(f"Task {name} not found") - # No-op if neither field was supplied. - if task_metadata is None and merge_params is None: + if task_metadata is None and merge_params is None and current_state is None: return task_entity - # `merge_params` is a separate atomic JSONB shallow-merge so concurrent - # callers don't overwrite each other's fields (vs reading→mutating→writing - # the whole params dict on task_entity). Run it first so the refreshed - # entity it returns becomes the base we apply `task_metadata` on top of; - # otherwise the `task_entity = merged` reassignment would discard an - # in-memory metadata change made before the merge. if merge_params: merged = await self.task_service.merge_task_params( task_entity.id, merge_params @@ -131,9 +125,18 @@ async def update_mutable_fields_on_task( if merged is not None: task_entity = merged + fields: dict[str, Any] = {} if task_metadata is not None: - task_entity.task_metadata = task_metadata - task_entity = await self.task_service.update_task(task=task_entity) + fields["task_metadata"] = task_metadata + if current_state is not None: + fields["current_state"] = current_state + if fields: + updated = await self.task_service.update_mutable_fields( + task_entity.id, fields + ) + if updated is None: + raise ItemDoesNotExist(f"Task {id or name} not found") + task_entity = updated return task_entity diff --git a/agentex/src/utils/task_constants.py b/agentex/src/utils/task_constants.py new file mode 100644 index 00000000..43d2b6bc --- /dev/null +++ b/agentex/src/utils/task_constants.py @@ -0,0 +1,2 @@ +# Single source for the current_state bound (request validation + column width). +CURRENT_STATE_MAX_LENGTH = 255 diff --git a/agentex/tests/integration/api/tasks/test_tasks_api.py b/agentex/tests/integration/api/tasks/test_tasks_api.py index dbf0a1a4..c14b700c 100644 --- a/agentex/tests/integration/api/tasks/test_tasks_api.py +++ b/agentex/tests/integration/api/tasks/test_tasks_api.py @@ -57,6 +57,18 @@ async def test_pagination_tasks(self, isolated_repositories, test_agent): tasks.append(await task_repo.create(agent_id=test_agent.id, task=task)) return tasks + @pytest_asyncio.fixture + async def test_running_task(self, isolated_repositories, test_agent): + """Create a minimal running task for tests that don't care about its name""" + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name=f"running-task-{orm_id()[:8]}", + status=TaskStatus.RUNNING, + status_reason="Running task for testing", + ) + return await task_repo.create(agent_id=test_agent.id, task=task) + @pytest_asyncio.fixture async def test_task_with_params(self, isolated_repositories, test_agent): """Create a test task with params directly via repository""" @@ -506,8 +518,7 @@ async def test_get_task_by_name_non_existent_returns_404(self, isolated_client): async def test_list_tasks_omits_params_in_response( self, isolated_client, test_task_with_params ): - """The list summary must omit `params` even when the task has them - (they can carry secrets/PII); fetch a single task for the full record.""" + """List summary omits params but carries current_state.""" # When - Request all tasks response = await isolated_client.get("/tasks") @@ -521,6 +532,8 @@ async def test_list_tasks_omits_params_in_response( ) assert params_task is not None, "Task should be in the list" assert "params" not in params_task + # current_state is a non-sensitive opaque label, so it IS in the lean summary. + assert "current_state" in params_task # async def test_get_task_by_id_includes_params_in_response( @@ -675,6 +688,102 @@ async def test_update_task_endpoint_success( assert response_data["task_metadata"]["configuration"]["version"] == "2.0.0" assert response_data["task_metadata"]["metrics"]["complexity_score"] == 75 + async def test_update_task_current_state( + self, isolated_client, test_running_task + ): + """PUT current_state: set it, omitted leaves it untouched, point-read reconciles.""" + created_task = test_running_task + + # Fresh task: current_state present in response and null by default. + response = await isolated_client.get(f"/tasks/{created_task.id}") + assert response.status_code == 200 + assert response.json()["current_state"] is None + + # Setting current_state persists and echoes back. + response = await isolated_client.put( + f"/tasks/{created_task.id}", json={"current_state": "awaiting_input"} + ) + assert response.status_code == 200 + assert response.json()["current_state"] == "awaiting_input" + + # Point-read reflects the committed value (source of truth). + response = await isolated_client.get(f"/tasks/{created_task.id}") + assert response.status_code == 200 + assert response.json()["current_state"] == "awaiting_input" + + # Updating only task_metadata (current_state omitted) does not clobber it. + response = await isolated_client.put( + f"/tasks/{created_task.id}", json={"task_metadata": {"k": "v"}} + ) + assert response.status_code == 200 + assert response.json()["current_state"] == "awaiting_input" + + async def test_update_task_current_state_and_metadata_together( + self, isolated_client, test_running_task + ): + """current_state + task_metadata in one PUT both persist without clobbering status.""" + created_task = test_running_task + + response = await isolated_client.put( + f"/tasks/{created_task.id}", + json={"current_state": "step_2", "task_metadata": {"stage": "two"}}, + ) + assert response.status_code == 200 + body = response.json() + assert body["current_state"] == "step_2" + assert body["task_metadata"] == {"stage": "two"} + assert body["status"] == "RUNNING" + + async def test_update_task_current_state_by_name( + self, isolated_client, isolated_repositories, test_agent + ): + """PUT /tasks/name/{name} forwards current_state too.""" + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-current-state-by-name", + status=TaskStatus.RUNNING, + status_reason="Test task for by-name current_state", + ) + await task_repo.create(agent_id=test_agent.id, task=task) + + response = await isolated_client.put( + "/tasks/name/task-for-current-state-by-name", + json={"current_state": "working"}, + ) + assert response.status_code == 200 + assert response.json()["current_state"] == "working" + + async def test_update_task_current_state_empty_string( + self, isolated_client, test_running_task + ): + """Empty string is a valid label distinct from null (guards a falsy-check regression).""" + response = await isolated_client.put( + f"/tasks/{test_running_task.id}", json={"current_state": ""} + ) + assert response.status_code == 200 + assert response.json()["current_state"] == "" + + async def test_update_task_current_state_too_long_rejected( + self, isolated_client, test_running_task + ): + """current_state exceeding the max length is rejected with 422.""" + response = await isolated_client.put( + f"/tasks/{test_running_task.id}", json={"current_state": "x" * 256} + ) + assert response.status_code == 422 + + async def test_update_task_request_ignores_unknown_fields( + self, isolated_client, test_running_task + ): + """Unknown fields are ignored (200, not 422) — guards the extra="ignore" SDK-compat assumption.""" + response = await isolated_client.put( + f"/tasks/{test_running_task.id}", + json={"current_state": "working", "field_from_a_newer_sdk": "ignored"}, + ) + assert response.status_code == 200 + assert response.json()["current_state"] == "working" + async def test_update_task_endpoint_validation( self, isolated_client, isolated_repositories ): diff --git a/agentex/tests/integration/test_task_stream.py b/agentex/tests/integration/test_task_stream.py index 918b0916..62109040 100644 --- a/agentex/tests/integration/test_task_stream.py +++ b/agentex/tests/integration/test_task_stream.py @@ -230,6 +230,63 @@ async def collect_stream_events(): print("✅ Task metadata update successfully triggered stream event") + async def test_current_state_update_triggers_stream_event( + self, test_agent_and_task, tasks_use_case, streams_use_case + ): + """current_state rides the existing task_updated event (reactive push to subscribers).""" + _agent, task = test_agent_and_task + + stream_events = [] + + async def collect_stream_events(): + try: + async for event_data in streams_use_case.stream_task_events( + task_id=task.id + ): + if event_data.startswith("data: "): + import json + + event_json = event_data[6:].strip() + if event_json: + try: + event = json.loads(event_json) + stream_events.append(event) + if event.get("type") == "task_updated": + break + except json.JSONDecodeError: + pass + except asyncio.CancelledError: + pass + + stream_task = asyncio.create_task(collect_stream_events()) + # Let the tail-only subscription establish before the update, or the event is missed. + await asyncio.sleep(0.1) + + updated_task = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="awaiting_input" + ) + + # Wait for the collector to see task_updated (it breaks on it); timeout is only a ceiling. + try: + async with asyncio.timeout(5): + await stream_task + except (TimeoutError, asyncio.CancelledError): + stream_task.cancel() + # Re-await so cancellation cleanup runs and no dangling-task warning leaks at teardown. + await asyncio.gather(stream_task, return_exceptions=True) + + task_updated_events = [ + e for e in stream_events if e.get("type") == "task_updated" + ] + assert len(task_updated_events) >= 1, ( + f"Expected task_updated event, got events: {[e.get('type') for e in stream_events]}" + ) + event_task = task_updated_events[0]["task"] + assert event_task["id"] == task.id + assert event_task["current_state"] == "awaiting_input" + + assert updated_task.current_state == "awaiting_input" + async def test_get_task_returns_updated_metadata_after_stream_update( self, test_agent_and_task, tasks_use_case ): diff --git a/agentex/tests/unit/services/test_task_service.py b/agentex/tests/unit/services/test_task_service.py index 91caa8d6..6267454e 100644 --- a/agentex/tests/unit/services/test_task_service.py +++ b/agentex/tests/unit/services/test_task_service.py @@ -938,6 +938,84 @@ async def test_update_task_with_task_metadata_changes( assert event_data["type"] == "task_updated" assert event_data["task"]["task_metadata"] == updated_metadata + async def test_update_task_current_state_publishes_stream_event( + self, task_service, agent_repository, sample_agent, redis_stream_repository + ): + """update_task persists current_state and carries it on the task_updated event.""" + await create_or_get_agent(agent_repository, sample_agent) + created_task = await task_service.create_task( + agent=sample_agent, task_name="task-for-current-state" + ) + + created_task.current_state = "working" + redis_stream_repository.send_data = AsyncMock() + + result = await task_service.update_task(created_task) + + assert result.current_state == "working" + retrieved_task = await task_service.get_task(id=created_task.id) + assert retrieved_task.current_state == "working" + + redis_stream_repository.send_data.assert_called_once() + call_args = redis_stream_repository.send_data.call_args + assert call_args[0][0] == f"task:{created_task.id}" + event_data = call_args[0][1] + assert event_data["type"] == "task_updated" + assert event_data["task"]["current_state"] == "working" + + async def test_update_mutable_fields_persists_and_publishes( + self, task_service, agent_repository, sample_agent, redis_stream_repository + ): + """update_mutable_fields persists the given columns and publishes task_updated.""" + await create_or_get_agent(agent_repository, sample_agent) + created_task = await task_service.create_task( + agent=sample_agent, task_name="task-for-mutable-fields" + ) + + redis_stream_repository.send_data = AsyncMock() + result = await task_service.update_mutable_fields( + created_task.id, + {"current_state": "working", "task_metadata": {"a": 1}}, + ) + + assert result.current_state == "working" + assert result.task_metadata == {"a": 1} + retrieved = await task_service.get_task(id=created_task.id) + assert retrieved.current_state == "working" + assert retrieved.task_metadata == {"a": 1} + + redis_stream_repository.send_data.assert_called_once() + event_data = redis_stream_repository.send_data.call_args[0][1] + assert event_data["type"] == "task_updated" + assert event_data["task"]["current_state"] == "working" + + async def test_update_mutable_fields_leaves_status_untouched( + self, task_service, agent_repository, sample_agent, redis_stream_repository + ): + """Writing current_state leaves status untouched (column-scoped update).""" + await create_or_get_agent(agent_repository, sample_agent) + created_task = await task_service.create_task( + agent=sample_agent, task_name="task-for-noclobber" + ) + + # Another writer moves the task to a terminal status. + await task_service.transition_task_status( + task_id=created_task.id, + expected_status=TaskStatus.RUNNING, + new_status=TaskStatus.COMPLETED, + status_reason="done", + ) + + redis_stream_repository.send_data = AsyncMock() + result = await task_service.update_mutable_fields( + created_task.id, {"current_state": "late"} + ) + + assert result.current_state == "late" + assert result.status == TaskStatus.COMPLETED + retrieved = await task_service.get_task(id=created_task.id) + assert retrieved.status == TaskStatus.COMPLETED + async def test_get_task_preserves_task_metadata( self, task_service, agent_repository, sample_agent ): diff --git a/agentex/tests/unit/use_cases/test_tasks_use_case.py b/agentex/tests/unit/use_cases/test_tasks_use_case.py index 798cb7cd..71fd95ab 100644 --- a/agentex/tests/unit/use_cases/test_tasks_use_case.py +++ b/agentex/tests/unit/use_cases/test_tasks_use_case.py @@ -3,6 +3,7 @@ methods (complete_task, fail_task, etc.) and metadata updates. """ +from unittest.mock import AsyncMock from uuid import uuid4 import pytest @@ -592,6 +593,122 @@ async def test_update_metadata_and_merge_params_both_persist( assert updated.task_metadata == {"stage": "tuned"} assert updated.params == {"model": "gpt-4", "temperature": 0.7} + async def test_update_current_state( + self, tasks_use_case, task_service, agent_repository, sample_agent + ): + """current_state persists and leaves status/task_metadata untouched.""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, + task_name="current-state-test", + task_metadata={"keep": "me"}, + ) + + updated = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="working" + ) + + assert updated.current_state == "working" + assert updated.status == TaskStatus.RUNNING + assert updated.task_metadata == {"keep": "me"} + + async def test_update_current_state_noop_when_omitted( + self, tasks_use_case, task_service, agent_repository, sample_agent + ): + """Omitting current_state leaves an already-set value untouched.""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, task_name="current-state-omitted-test" + ) + await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="set-once" + ) + + # A later metadata-only update must not clear current_state. + updated = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, task_metadata={"a": 1} + ) + + assert updated.current_state == "set-once" + + async def test_update_current_state_and_metadata_single_atomic_write( + self, tasks_use_case, task_service, agent_repository, sample_agent + ): + """current_state + task_metadata together persist via one atomic write (one publish).""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, task_name="current-state-combined-test" + ) + + spy = AsyncMock(wraps=task_service.update_mutable_fields) + task_service.update_mutable_fields = spy + + updated = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, + task_metadata={"stage": "two"}, + current_state="working", + ) + + spy.assert_awaited_once() + assert updated.current_state == "working" + assert updated.task_metadata == {"stage": "two"} + assert updated.status == TaskStatus.RUNNING + + async def test_update_current_state_does_not_clobber_concurrent_status( + self, + tasks_use_case, + task_service, + task_repository, + agent_repository, + sample_agent, + ): + """Stale-read race: current_state write must not revert a concurrent status transition.""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, task_name="current-state-clobber-test" + ) + + # Snapshot the entity BEFORE the transition — the stale read the old bug wrote back. + stale_entity = await task_service.get_task(id=task.id) + assert stale_entity.status == TaskStatus.RUNNING + + # Another writer moves the task to a terminal status after that read. + await task_service.transition_task_status( + task_id=task.id, + expected_status=TaskStatus.RUNNING, + new_status=TaskStatus.COMPLETED, + status_reason="done", + ) + + # Force the use case to operate on the stale (pre-transition) read. + task_service.get_task = AsyncMock(return_value=stale_entity) + + updated = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="late" + ) + + # The write set current_state without reverting the terminal status. + assert updated.current_state == "late" + assert updated.status == TaskStatus.COMPLETED + persisted = await task_repository.get(id=task.id) + assert persisted.status == TaskStatus.COMPLETED + assert persisted.current_state == "late" + + async def test_update_current_state_on_deleted_task_raises( + self, tasks_use_case, task_service, agent_repository, sample_agent + ): + """Updating current_state on a deleted task raises not found.""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, task_name="current-state-deleted-test" + ) + await tasks_use_case.delete_task(id=task.id) + + with pytest.raises(ItemDoesNotExist): + await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="working" + ) + async def test_update_metadata_on_deleted_task_raises( self, tasks_use_case, task_service, agent_repository, sample_agent ):