Skip to content

Commit bfd7554

Browse files
authored
Add order by and order direction field to mongodb stores (#85)
Adds the order by and order direction fields to the messages and states objects, along with some tests.
1 parent 3570792 commit bfd7554

9 files changed

Lines changed: 345 additions & 31 deletions

File tree

agentex/src/adapters/crud_store/adapter_mongodb.py

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -510,6 +510,8 @@ async def list(
510510
filters: dict[str, Any] | None = None,
511511
limit: int | None = None,
512512
page_number: int | None = None,
513+
order_by: str | None = None,
514+
order_direction: str | None = None,
513515
) -> list[T]:
514516
"""
515517
List all documents in the collection.
@@ -525,8 +527,20 @@ async def list(
525527
else:
526528
cursor = self.collection.find()
527529
cursor = cursor.skip(skip).limit(limit)
528-
# Consistent sorting so that pagination works correctly
529-
cursor = cursor.sort([("_id", 1)])
530+
531+
sort_list = []
532+
if order_by:
533+
direction = (
534+
pymongo.DESCENDING
535+
if order_direction and order_direction.lower() == "desc"
536+
else pymongo.ASCENDING
537+
)
538+
sort_list.append((order_by, direction))
539+
540+
# Always use _id as tiebreaker
541+
sort_list.append(("_id", pymongo.ASCENDING))
542+
cursor = cursor.sort(sort_list)
543+
530544
try:
531545
return [self._deserialize(doc) for doc in cursor]
532546
except Exception as e:

agentex/src/api/routes/messages.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -115,9 +115,15 @@ async def list_messages(
115115
message_use_case: DMessageUseCase,
116116
limit: int = 50,
117117
page_number: int = 1,
118+
order_by: str | None = None,
119+
order_direction: str = "desc",
118120
) -> list[TaskMessage]:
119121
task_message_entities = await message_use_case.list_messages(
120-
task_id=task_id, limit=limit, page_number=page_number
122+
task_id=task_id,
123+
limit=limit,
124+
page_number=page_number,
125+
order_by=order_by,
126+
order_direction=order_direction,
121127
)
122128

123129
return [

agentex/src/api/routes/states.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -62,9 +62,16 @@ async def filter_states(
6262
agent_id: str | None = Query(None, description="Agent ID"),
6363
limit: int = Query(50, description="Limit", ge=1),
6464
page_number: int = Query(1, description="Page number", ge=1),
65+
order_by: str | None = Query(None, description="Field to order by"),
66+
order_direction: str = Query("desc", description="Order direction (asc or desc)"),
6567
) -> list[State]:
6668
state_entities = await states_use_case.list(
67-
task_id=task_id, agent_id=agent_id, limit=limit, page_number=page_number
69+
task_id=task_id,
70+
agent_id=agent_id,
71+
limit=limit,
72+
page_number=page_number,
73+
order_by=order_by,
74+
order_direction=order_direction,
6875
)
6976
logger.info(f"Listing states: {state_entities}")
7077
return [State.model_validate(state_entity) for state_entity in state_entities]

agentex/src/domain/services/task_message_service.py

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -43,26 +43,35 @@ async def get_message(self, message_id: str) -> TaskMessageEntity:
4343
return await self.repository.get(id=message_id)
4444

4545
async def get_messages(
46-
self, task_id: str, limit: int, page_number: int
46+
self,
47+
task_id: str,
48+
limit: int,
49+
page_number: int,
50+
order_by: str | None = None,
51+
order_direction: str = "desc",
4752
) -> list[TaskMessageEntity]:
4853
"""
4954
Get all messages for a specific task.
5055
5156
Args:
5257
task_id: The task ID
5358
limit: Optional limit on the number of messages to return
59+
order_by: Optional field name to order by (defaults to created_at)
60+
order_direction: Optional direction to order by ("asc" or "desc", defaults to "desc")
5461
5562
Returns:
5663
List of TaskMessageEntity objects for the task
5764
"""
58-
# Sort by created_at in ascending order (oldest first)
59-
# This is typically what we want for conversation history
65+
# Default to created_at descending (newest first)
66+
sort_field = order_by or "created_at"
67+
sort_direction = 1 if order_direction.lower() == "asc" else -1
68+
6069
return await self.repository.find_by_field(
6170
"task_id",
6271
task_id,
6372
limit=limit,
6473
page_number=page_number,
65-
sort_by={"created_at": 1},
74+
sort_by={sort_field: sort_direction},
6675
)
6776

6877
async def append_message(

agentex/src/domain/use_cases/messages_use_case.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -105,20 +105,31 @@ async def get_message(self, message_id: str) -> TaskMessageEntity | None:
105105
return await self.task_message_service.get_message(message_id=message_id)
106106

107107
async def list_messages(
108-
self, task_id: str, limit: int, page_number: int
108+
self,
109+
task_id: str,
110+
limit: int,
111+
page_number: int,
112+
order_by: str | None = None,
113+
order_direction: str = "desc",
109114
) -> list[TaskMessageEntity]:
110115
"""
111116
Get all messages for a task.
112117
113118
Args:
114119
task_id: The task ID
115120
limit: Optional limit on the number of messages to return
121+
order_by: Optional field name to order by (defaults to created_at)
122+
order_direction: Optional direction to order by ("asc" or "desc", defaults to "desc")
116123
117124
Returns:
118125
List of TaskMessageEntity objects for the task
119126
"""
120127
return await self.task_message_service.get_messages(
121-
task_id=task_id, limit=limit, page_number=page_number
128+
task_id=task_id,
129+
limit=limit,
130+
page_number=page_number,
131+
order_by=order_by,
132+
order_direction=order_direction,
122133
)
123134

124135

agentex/src/domain/use_cases/states_use_case.py

Lines changed: 15 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -35,25 +35,22 @@ async def list(
3535
page_number: int,
3636
task_id: str | None = None,
3737
agent_id: str | None = None,
38+
order_by: str | None = None,
39+
order_direction: str = "desc",
3840
) -> list[StateEntity]:
39-
if task_id and agent_id:
40-
return await self.task_state_repository.list(
41-
filters={"task_id": task_id, "agent_id": agent_id},
42-
limit=limit,
43-
page_number=page_number,
44-
)
45-
elif task_id:
46-
return await self.task_state_repository.list(
47-
filters={"task_id": task_id}, limit=limit, page_number=page_number
48-
)
49-
elif agent_id:
50-
return await self.task_state_repository.list(
51-
filters={"agent_id": agent_id}, limit=limit, page_number=page_number
52-
)
53-
else:
54-
return await self.task_state_repository.list(
55-
limit=limit, page_number=page_number
56-
)
41+
filters = {}
42+
if task_id:
43+
filters["task_id"] = task_id
44+
if agent_id:
45+
filters["agent_id"] = agent_id
46+
47+
return await self.task_state_repository.list(
48+
filters=filters if filters else None,
49+
limit=limit,
50+
page_number=page_number,
51+
order_by=order_by,
52+
order_direction=order_direction,
53+
)
5754

5855
async def update(self, id: str, task_id: str, state: dict[str, Any]) -> StateEntity:
5956
task_state = await self.task_state_repository.get(id=id)

agentex/tests/integration/api/messages/test_messages_api.py

Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -289,3 +289,138 @@ async def test_list_messages_pagination(
289289
assert {(d["id"], d["content"]["content"]) for d in paginated_messages} == {
290290
(d.id, d.content.content) for d in test_pagination_messages
291291
}
292+
293+
async def test_list_messages_with_order_by(
294+
self, isolated_client, isolated_repositories
295+
):
296+
"""Test that list messages endpoint supports order_by parameter"""
297+
# Given - Create an agent and task
298+
agent_repo = isolated_repositories["agent_repository"]
299+
agent = AgentEntity(
300+
id=orm_id(),
301+
name="order-by-message-agent",
302+
description="Agent for order_by message testing",
303+
acp_url="http://test-acp:8000",
304+
acp_type=ACPType.SYNC,
305+
)
306+
await agent_repo.create(agent)
307+
308+
task_repo = isolated_repositories["task_repository"]
309+
task = TaskEntity(
310+
id=orm_id(),
311+
name="order-by-message-task",
312+
status=TaskStatus.RUNNING,
313+
status_reason="Task for order_by message testing",
314+
)
315+
await task_repo.create(agent_id=agent.id, task=task)
316+
317+
# Create multiple messages
318+
message_repo = isolated_repositories["task_message_repository"]
319+
messages = []
320+
for i in range(3):
321+
message = TaskMessageEntity(
322+
id=orm_id(),
323+
task_id=task.id,
324+
content=TextContentEntity(
325+
type="text", author="user", content=f"Order test message {i}"
326+
),
327+
streaming_status="DONE",
328+
)
329+
messages.append(await message_repo.create(message))
330+
331+
# When - Request messages with order_by=created_at and order_direction=asc
332+
response_asc = await isolated_client.get(
333+
"/messages",
334+
params={
335+
"task_id": task.id,
336+
"order_by": "created_at",
337+
"order_direction": "asc",
338+
},
339+
)
340+
341+
# Then - Should return messages in ascending order
342+
assert response_asc.status_code == 200
343+
messages_asc = response_asc.json()
344+
assert len(messages_asc) == 3
345+
346+
# Verify ascending order
347+
for i in range(len(messages_asc) - 1):
348+
assert messages_asc[i]["created_at"] <= messages_asc[i + 1]["created_at"]
349+
350+
# When - Request messages with order_by=created_at and order_direction=desc
351+
response_desc = await isolated_client.get(
352+
"/messages",
353+
params={
354+
"task_id": task.id,
355+
"order_by": "created_at",
356+
"order_direction": "desc",
357+
},
358+
)
359+
360+
# Then - Should return messages in descending order
361+
assert response_desc.status_code == 200
362+
messages_desc = response_desc.json()
363+
assert len(messages_desc) == 3
364+
365+
# Verify descending order
366+
for i in range(len(messages_desc) - 1):
367+
assert messages_desc[i]["created_at"] >= messages_desc[i + 1]["created_at"]
368+
369+
# Verify asc and desc return different orderings (first element of asc should be in last position of desc)
370+
# Note: We only check the first element reversal since items with identical timestamps
371+
# may have unpredictable relative ordering
372+
assert messages_asc[0]["id"] == messages_desc[-1]["id"]
373+
374+
async def test_list_messages_order_by_defaults_to_desc(
375+
self, isolated_client, isolated_repositories
376+
):
377+
"""Test that order_direction defaults to desc for messages (newest first)"""
378+
# Given - Create an agent and task
379+
agent_repo = isolated_repositories["agent_repository"]
380+
agent = AgentEntity(
381+
id=orm_id(),
382+
name="order-default-message-agent",
383+
description="Agent for order default message testing",
384+
acp_url="http://test-acp:8000",
385+
acp_type=ACPType.SYNC,
386+
)
387+
await agent_repo.create(agent)
388+
389+
task_repo = isolated_repositories["task_repository"]
390+
task = TaskEntity(
391+
id=orm_id(),
392+
name="order-default-message-task",
393+
status=TaskStatus.RUNNING,
394+
status_reason="Task for order default message testing",
395+
)
396+
await task_repo.create(agent_id=agent.id, task=task)
397+
398+
# Create multiple messages
399+
message_repo = isolated_repositories["task_message_repository"]
400+
for i in range(3):
401+
message = TaskMessageEntity(
402+
id=orm_id(),
403+
task_id=task.id,
404+
content=TextContentEntity(
405+
type="text", author="user", content=f"Default order message {i}"
406+
),
407+
streaming_status="DONE",
408+
)
409+
await message_repo.create(message)
410+
411+
# When - Request messages without specifying order_direction
412+
response = await isolated_client.get(
413+
"/messages",
414+
params={"task_id": task.id},
415+
)
416+
417+
# Then - Should return messages successfully
418+
assert response.status_code == 200
419+
messages = response.json()
420+
assert len(messages) == 3
421+
422+
# Verify descending order - items with same timestamp may have any relative order,
423+
# so we only check that timestamps are non-increasing (allowing equal timestamps)
424+
timestamps = [m["created_at"] for m in messages]
425+
# Sort timestamps descending and verify the returned order matches a valid descending order
426+
assert timestamps == sorted(timestamps, reverse=True)

0 commit comments

Comments
 (0)