diff --git a/CLAUDE.md b/CLAUDE.md index 99a62276..f8530835 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -358,7 +358,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 | |---|---| @@ -371,7 +371,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 904f1615..056cc3dc 100644 --- a/agentex/openapi.yaml +++ b/agentex/openapi.yaml @@ -6588,6 +6588,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'. type: object required: - id @@ -6811,6 +6817,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: @@ -7396,6 +7408,13 @@ 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; pass null to + clear it, omit to leave it unchanged. 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 86abb22a..b87503d3 100644 --- a/agentex/src/api/routes/tasks.py +++ b/agentex/src/api/routes/tasks.py @@ -22,7 +22,7 @@ from src.domain.entities.tasks import TaskStatus as DomainTaskStatus from src.domain.services.authorization_service import DAuthorizationService from src.domain.use_cases.streams_use_case import DStreamsUseCase -from src.domain.use_cases.tasks_use_case import DTaskUseCase +from src.domain.use_cases.tasks_use_case import UNSET, DTaskUseCase from src.utils.authorization_shortcuts import ( DAuthorizedId, DAuthorizedName, @@ -196,6 +196,12 @@ async def update_task( id=task_id, task_metadata=request.task_metadata, merge_params=request.merge_params, + # Forward only when the client sent the key: explicit null clears, omitted leaves unchanged. + current_state=( + request.current_state + if "current_state" in request.model_fields_set + else UNSET + ), ) return Task.model_validate(updated_task_entity) @@ -217,6 +223,12 @@ async def update_task_by_name( name=task_name, task_metadata=request.task_metadata, merge_params=request.merge_params, + # Forward only when the client sent the key: explicit null clears, omitted leaves unchanged. + current_state=( + request.current_state + if "current_state" in request.model_fields_set + else UNSET + ), ) return Task.model_validate(updated_task_entity) diff --git a/agentex/src/api/schemas/tasks.py b/agentex/src/api/schemas/tasks.py index a79d9fef..5b53f6e1 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): @@ -63,6 +64,14 @@ class Task(BaseModel): None, title="Task metadata", ) + # No bound on the read path: enforcing it would 500 if the column is ever widened (writes are already bounded). + 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'." + ), + ) class TaskResponse(Task): @@ -87,6 +96,14 @@ 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; " + "pass null to clear it, omit to leave it unchanged." + ), + ) class TaskStatusReasonRequest(BaseModel): diff --git a/agentex/src/domain/entities/tasks.py b/agentex/src/domain/entities/tasks.py index 76e93cf9..82cc3776 100644 --- a/agentex/src/domain/entities/tasks.py +++ b/agentex/src/domain/entities/tasks.py @@ -66,6 +66,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") @@ -84,4 +91,5 @@ def convert_task_to_entity(task: Task) -> TaskEntity: cleaned_at=task.cleaned_at, params=task.params, task_metadata=task.task_metadata, + current_state=task.current_state, ) diff --git a/agentex/src/domain/repositories/task_repository.py b/agentex/src/domain/repositories/task_repository.py index 45662015..439a6802 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,10 @@ logger = make_logger(__name__) +# Columns update_mutable_fields is allowed to set — a code guardrail (not a comment) +# so a future caller can't route status/params through it and bypass their atomic CAS. +_MUTABLE_TASK_COLUMNS = frozenset({"task_metadata", "current_state"}) + class TaskRepository(PostgresCRUDRepository[TaskORM, TaskEntity, TaskRelationships]): """Repository for Task entity with relationship loading support""" @@ -231,20 +235,42 @@ 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: + """Atomically set the given caller-mutable columns on one task row via + ``UPDATE ... RETURNING`` — column-scoped, so (unlike ``update``'s whole-row merge) it + can't clobber a concurrently changed ``status``/``params``. Values apply verbatim + (``current_state=None`` clears). Returns the updated entity, or ``None`` if absent. + """ + 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 tasks SET WHERE id`` → the updated entity, or ``None`` if no + such task exists. 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 d5d1291f..59e2e6cf 100644 --- a/agentex/src/domain/services/task_service.py +++ b/agentex/src/domain/services/task_service.py @@ -242,6 +242,32 @@ async def update_task(self, task: TaskEntity) -> TaskEntity: return updated_task + async def update_mutable_fields( + self, task_id: str, fields: dict[str, Any] + ) -> TaskEntity | None: + """Column-scoped atomic update of the given columns, then publish task_updated. + Returns the updated entity, or ``None`` if the task no longer exists. + """ + updated_task = await self.task_repository.update_mutable_fields(task_id, fields) + 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 + ) + + 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. @@ -374,7 +400,11 @@ async def cancel_task( new_status=TaskStatus.CANCELED, status_reason="Task canceled by user", ) - return updated if updated is not None else await self.task_repository.get(id=task.id) + return ( + updated + if updated is not None + else await self.task_repository.get(id=task.id) + ) async def interrupt_task( self, agent: AgentEntity, task: TaskEntity, acp_url: str diff --git a/agentex/src/domain/use_cases/tasks_use_case.py b/agentex/src/domain/use_cases/tasks_use_case.py index 7c114744..3812b348 100644 --- a/agentex/src/domain/use_cases/tasks_use_case.py +++ b/agentex/src/domain/use_cases/tasks_use_case.py @@ -11,6 +11,13 @@ logger = make_logger(__name__) +class _Unset: + """Sentinel for partial updates: field omitted (untouched) vs. explicitly null (cleared).""" + + +UNSET: Any = _Unset() + + class TasksUseCase: """ Use case for managing tasks. Handles CRUD operations and delegates task operations to ACP servers. @@ -95,12 +102,17 @@ 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 | _Unset = UNSET, ) -> 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. ``current_state`` uses the UNSET sentinel + (explicit null clears, omitted leaves it); ``task_metadata`` None means "not supplied". + """ if not id and not name: raise ClientError("Either id or name must be provided") + current_state_supplied = not isinstance(current_state, _Unset) + # todo: make this a transaction? task_entity = await self.task_service.get_task(id=id, name=name) if task_entity.status == TaskStatus.DELETED: @@ -109,16 +121,15 @@ 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: + # No-op if no mutable field was supplied. + if ( + task_metadata is None + and merge_params is None + and not current_state_supplied + ): 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. + # Atomic JSONB shallow-merge; run first so its refreshed entity is the fallback return. if merge_params: merged = await self.task_service.merge_task_params( task_entity.id, merge_params @@ -126,9 +137,20 @@ async def update_mutable_fields_on_task( if merged is not None: task_entity = merged + # Single column-scoped write → one task_updated publish, no whole-row clobber. + 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_supplied: + fields["current_state"] = current_state + if fields: + updated = await self.task_service.update_mutable_fields( + task_entity.id, fields + ) + if updated is None: + # Row vanished mid-flight (defensive; no live hard-delete path). Raise, don't return stale. + 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..d1f0b46f --- /dev/null +++ b/agentex/src/utils/task_constants.py @@ -0,0 +1,3 @@ +# Single source for the current_state bound: the request-validation limit +# (UpdateTaskRequest) and the storage limit (TaskORM.current_state 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 bdf857bd..f292434d 100644 --- a/agentex/tests/integration/api/tasks/test_tasks_api.py +++ b/agentex/tests/integration/api/tasks/test_tasks_api.py @@ -683,6 +683,212 @@ 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, isolated_repositories + ): + """PUT current_state: explicit null clears, omitted leaves it, point-read reconciles.""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="current-state-agent", + description="Agent for current_state update testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-current-state", + status=TaskStatus.RUNNING, + status_reason="Test task for current_state", + ) + created_task = await task_repo.create(agent_id=agent.id, task=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" + + # Explicit null clears the label (distinct from omitting the field). + response = await isolated_client.put( + f"/tasks/{created_task.id}", json={"current_state": None} + ) + assert response.status_code == 200 + assert response.json()["current_state"] is None + response = await isolated_client.get(f"/tasks/{created_task.id}") + assert response.json()["current_state"] is None + + async def test_update_task_current_state_and_metadata_together( + self, isolated_client, isolated_repositories + ): + """current_state + task_metadata in one PUT both persist without clobbering status.""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="current-state-combined-agent", + description="Agent for combined update testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-combined-update", + status=TaskStatus.RUNNING, + status_reason="Test task for combined update", + ) + created_task = await task_repo.create(agent_id=agent.id, task=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 + ): + """PUT /tasks/name/{name} forwards current_state too.""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="current-state-by-name-agent", + description="Agent for by-name current_state testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + 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=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, isolated_repositories + ): + """Empty string is a valid label distinct from null (guards a falsy-check regression).""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="current-state-empty-agent", + description="Agent for empty-string current_state testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-current-state-empty", + status=TaskStatus.RUNNING, + status_reason="Test task for empty current_state", + ) + created_task = await task_repo.create(agent_id=agent.id, task=task) + + response = await isolated_client.put( + f"/tasks/{created_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, isolated_repositories + ): + """current_state exceeding the max length is rejected with 422.""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="current-state-toolong-agent", + description="Agent for max-length current_state testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-current-state-toolong", + status=TaskStatus.RUNNING, + status_reason="Test task for over-long current_state", + ) + created_task = await task_repo.create(agent_id=agent.id, task=task) + + response = await isolated_client.put( + f"/tasks/{created_task.id}", json={"current_state": "x" * 256} + ) + assert response.status_code == 422 + + async def test_update_task_request_ignores_unknown_fields( + self, isolated_client, isolated_repositories + ): + """Unknown fields are ignored (200, not 422) — guards the extra="ignore" SDK-compat assumption.""" + agent_repo = isolated_repositories["agent_repository"] + agent = AgentEntity( + id=orm_id(), + name="unknown-fields-agent", + description="Agent for unknown-field compat testing", + acp_url="http://test-acp:8000", + acp_type=ACPType.SYNC, + ) + await agent_repo.create(agent) + + task_repo = isolated_repositories["task_repository"] + task = TaskEntity( + id=orm_id(), + name="task-for-unknown-fields", + status=TaskStatus.RUNNING, + status_reason="Test task for unknown-field compat", + ) + created_task = await task_repo.create(agent_id=agent.id, task=task) + + response = await isolated_client.put( + f"/tasks/{created_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 7da23edd..ae0786a6 100644 --- a/agentex/tests/integration/test_task_stream.py +++ b/agentex/tests/integration/test_task_stream.py @@ -230,6 +230,65 @@ 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" + + print("✅ current_state update successfully triggered stream event") + 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..68f4865e 100644 --- a/agentex/tests/unit/services/test_task_service.py +++ b/agentex/tests/unit/services/test_task_service.py @@ -938,6 +938,86 @@ 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 + ): + """The primitive is column-scoped: writing current_state leaves status untouched + (the use-case stale-read clobber regression is guarded in test_tasks_use_case.py). + """ + 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..bc523076 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,144 @@ 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 (the UNSET default) leaves it 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_clears_on_explicit_null( + self, tasks_use_case, task_service, agent_repository, sample_agent + ): + """Passing current_state=None explicitly clears the label (vs omitting).""" + await create_or_get_agent(agent_repository, sample_agent) + task = await task_service.create_task( + agent=sample_agent, task_name="current-state-clear-test" + ) + was_set = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state="working" + ) + # Confirm it was set, so the clear below is a real transition (not a trivial null→null). + assert was_set.current_state == "working" + + cleared = await tasks_use_case.update_mutable_fields_on_task( + id=task.id, current_state=None + ) + + assert cleared.current_state is None + + 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, + ): + """Regression: a stale RUNNING read racing a COMPLETED transition must not revert + status on the current_state write. Fails on the old whole-row merge, passes column-scoped. + """ + 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 ):