Skip to content

Commit fab413d

Browse files
committed
feat(sessions): start-of-turn ledger rows with completion, and turn-index ordering
1 parent d9da538 commit fab413d

13 files changed

Lines changed: 752 additions & 77 deletions

File tree

api/oss/src/apis/fastapi/sessions/models.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,13 @@ class SessionTurnAppendRequest(BaseModel):
175175
end_time: Optional[datetime] = None
176176

177177

178+
class SessionTurnCompleteRequest(BaseModel):
179+
session_id: str
180+
turn_index: int
181+
agent_session_id: Optional[str] = None
182+
end_time: datetime
183+
184+
178185
class SessionTurnQueryRequest(BaseModel):
179186
query: Optional[SessionTurnQuery] = None
180187
windowing: Optional[Windowing] = None

api/oss/src/apis/fastapi/sessions/router.py

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -61,8 +61,9 @@
6161
from oss.src.core.sessions.interactions.types import InteractionNotFound
6262
from oss.src.core.sessions.mounts.service import SessionMountsService
6363
from oss.src.core.sessions.mounts.dtos import SessionMountQuery
64-
from oss.src.core.sessions.turns.dtos import SessionTurnCreate
64+
from oss.src.core.sessions.turns.dtos import SessionTurnComplete, SessionTurnCreate
6565
from oss.src.core.sessions.turns.service import SessionTurnsService
66+
from oss.src.core.sessions.turns.types import SessionTurnNotFound
6667
from oss.src.core.sessions.dtos import SessionQuery
6768
from oss.src.core.sessions.service import SessionsService
6869
from oss.src.core.mounts.service import MountsService
@@ -107,6 +108,7 @@
107108
SessionMountsResponse,
108109
# turns
109110
SessionTurnAppendRequest,
111+
SessionTurnCompleteRequest,
110112
SessionTurnQueryRequest,
111113
SessionTurnResponse,
112114
SessionTurnsResponse,
@@ -1048,6 +1050,15 @@ def __init__(self, *, turns_service: SessionTurnsService) -> None:
10481050
response_model=SessionTurnResponse,
10491051
response_model_exclude_none=True,
10501052
)
1053+
self.router.add_api_route(
1054+
"/complete",
1055+
self.complete_turn,
1056+
methods=["POST"],
1057+
operation_id="complete_turn",
1058+
status_code=status.HTTP_200_OK,
1059+
response_model=SessionTurnResponse,
1060+
response_model_exclude_none=True,
1061+
)
10511062
self.router.add_api_route(
10521063
"/query",
10531064
self.query_turns,
@@ -1102,6 +1113,39 @@ async def append_turn(
11021113
)
11031114
return SessionTurnResponse(count=1, turn=turn)
11041115

1116+
@intercept_exceptions()
1117+
async def complete_turn(
1118+
self,
1119+
request: Request,
1120+
body: SessionTurnCompleteRequest,
1121+
) -> SessionTurnResponse:
1122+
project_id: UUID = request.state.project_id
1123+
user_id: UUID = request.state.user_id
1124+
1125+
if not await check_action_access(
1126+
user_uid=str(user_id),
1127+
project_id=str(project_id),
1128+
permission=Permission.RUN_SESSIONS,
1129+
):
1130+
raise FORBIDDEN_EXCEPTION
1131+
1132+
try:
1133+
turn = await self.turns_service.complete_turn(
1134+
project_id=project_id,
1135+
turn=SessionTurnComplete(
1136+
session_id=body.session_id,
1137+
turn_index=body.turn_index,
1138+
agent_session_id=body.agent_session_id,
1139+
end_time=body.end_time,
1140+
),
1141+
)
1142+
except SessionTurnNotFound as e:
1143+
raise HTTPException(
1144+
status_code=status.HTTP_404_NOT_FOUND,
1145+
detail=e.message,
1146+
) from e
1147+
return SessionTurnResponse(count=1, turn=turn)
1148+
11051149
@intercept_exceptions()
11061150
async def query_turns(
11071151
self,

api/oss/src/core/sessions/turns/dtos.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,13 @@ class SessionTurnCreate(BaseModel):
3838
end_time: Optional[datetime] = None
3939

4040

41+
class SessionTurnComplete(BaseModel):
42+
session_id: str
43+
turn_index: int
44+
agent_session_id: Optional[str] = None
45+
end_time: datetime
46+
47+
4148
class SessionTurnQuery(BaseModel):
4249
session_id: Optional[str] = None
4350
stream_id: Optional[UUID] = None

api/oss/src/core/sessions/turns/interfaces.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
from oss.src.core.sessions.turns.dtos import (
77
HarnessKind,
88
SessionTurn,
9+
SessionTurnComplete,
910
SessionTurnCreate,
1011
SessionTurnQuery,
1112
)
@@ -22,6 +23,15 @@ async def append(
2223
turn: SessionTurnCreate,
2324
) -> SessionTurn: ...
2425

26+
@abstractmethod
27+
async def complete(
28+
self,
29+
*,
30+
project_id: UUID,
31+
#
32+
turn: SessionTurnComplete,
33+
) -> Optional[SessionTurn]: ...
34+
2535
@abstractmethod
2636
async def fetch_turn(
2737
self,

api/oss/src/core/sessions/turns/service.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,10 +12,12 @@
1212
from oss.src.core.sessions.turns.dtos import (
1313
HarnessKind,
1414
SessionTurn,
15+
SessionTurnComplete,
1516
SessionTurnCreate,
1617
SessionTurnQuery,
1718
)
1819
from oss.src.core.sessions.turns.interfaces import SessionTurnsDAOInterface
20+
from oss.src.core.sessions.turns.types import SessionTurnNotFound
1921

2022

2123
class SessionTurnsService:
@@ -52,6 +54,21 @@ async def fetch_turn(
5254
turn_id=turn_id,
5355
)
5456

57+
async def complete_turn(
58+
self,
59+
*,
60+
project_id: UUID,
61+
#
62+
turn: SessionTurnComplete,
63+
) -> SessionTurn:
64+
completed = await self._dao.complete(
65+
project_id=project_id,
66+
turn=turn,
67+
)
68+
if completed is None:
69+
raise SessionTurnNotFound(turn.session_id, turn.turn_index)
70+
return completed
71+
5572
async def query_turns(
5673
self,
5774
*,

api/oss/src/core/sessions/turns/types.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,8 @@ class SessionTurnError(Exception):
66

77

88
class SessionTurnNotFound(SessionTurnError):
9-
def __init__(self, turn_id: str):
10-
self.turn_id = turn_id
11-
self.message = f"No turn found with id '{turn_id}'."
9+
def __init__(self, session_id: str, turn_index: int):
10+
self.session_id = session_id
11+
self.turn_index = turn_index
12+
self.message = f"No turn {turn_index} found for session '{session_id}'."
1213
super().__init__(self.message)

api/oss/src/dbs/postgres/sessions/turns/dao.py

Lines changed: 103 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,14 @@
11
from typing import List, Optional
22
from uuid import UUID
33

4-
from sqlalchemy import delete as sa_delete, select
4+
from sqlalchemy import and_, delete as sa_delete, or_, select
5+
from sqlalchemy import update as sa_update
56
from sqlalchemy.exc import IntegrityError
67

78
from oss.src.core.sessions.turns.dtos import (
89
HarnessKind,
910
SessionTurn,
11+
SessionTurnComplete,
1012
SessionTurnCreate,
1113
SessionTurnQuery,
1214
)
@@ -23,7 +25,6 @@
2325
TransactionsEngine,
2426
get_transactions_engine,
2527
)
26-
from oss.src.dbs.postgres.shared.utils import apply_windowing
2728

2829

2930
class SessionTurnsDAO(SessionTurnsDAOInterface):
@@ -68,6 +69,48 @@ async def append(
6869
raise
6970
return map_turn_dbe_to_dto(turn_dbe=dbe)
7071

72+
async def complete(
73+
self,
74+
*,
75+
project_id: UUID,
76+
#
77+
turn: SessionTurnComplete,
78+
) -> Optional[SessionTurn]:
79+
values = {"end_time": turn.end_time}
80+
if turn.agent_session_id is not None:
81+
values["agent_session_id"] = turn.agent_session_id
82+
83+
async with self.engine.session() as session:
84+
stmt = (
85+
sa_update(SessionTurnDBE)
86+
.where(
87+
SessionTurnDBE.project_id == project_id,
88+
SessionTurnDBE.session_id == turn.session_id,
89+
SessionTurnDBE.turn_index == turn.turn_index,
90+
SessionTurnDBE.end_time.is_(None),
91+
)
92+
.values(**values)
93+
.returning(SessionTurnDBE)
94+
)
95+
result = await session.execute(stmt)
96+
dbe = result.scalar_one_or_none()
97+
if dbe is not None:
98+
await session.commit()
99+
await session.refresh(dbe)
100+
return map_turn_dbe_to_dto(turn_dbe=dbe)
101+
102+
stmt = select(SessionTurnDBE).where(
103+
SessionTurnDBE.project_id == project_id,
104+
SessionTurnDBE.session_id == turn.session_id,
105+
SessionTurnDBE.turn_index == turn.turn_index,
106+
)
107+
result = await session.execute(stmt)
108+
dbe = result.scalar_one_or_none()
109+
110+
if dbe is None:
111+
return None
112+
return map_turn_dbe_to_dto(turn_dbe=dbe)
113+
71114
async def fetch_turn(
72115
self,
73116
*,
@@ -119,16 +162,67 @@ async def query_turns(
119162
SessionTurnDBE.references.contains(turn_references),
120163
)
121164

165+
descending = windowing is None or windowing.order != "ascending"
166+
time_attribute = SessionTurnDBE.start_time
167+
122168
if windowing:
123-
stmt = apply_windowing(
124-
stmt=stmt,
125-
DBE=SessionTurnDBE,
126-
attribute="id",
127-
order="descending",
128-
windowing=windowing,
169+
if windowing.newest is not None:
170+
if descending and windowing.next is None:
171+
stmt = stmt.where(time_attribute < windowing.newest)
172+
else:
173+
stmt = stmt.where(time_attribute <= windowing.newest)
174+
if windowing.oldest is not None:
175+
if not descending and windowing.next is None:
176+
stmt = stmt.where(time_attribute > windowing.oldest)
177+
else:
178+
stmt = stmt.where(time_attribute >= windowing.oldest)
179+
180+
if windowing.next is not None:
181+
cursor_stmt = select(
182+
SessionTurnDBE.turn_index,
183+
SessionTurnDBE.id,
184+
).where(
185+
SessionTurnDBE.project_id == project_id,
186+
SessionTurnDBE.id == windowing.next,
187+
)
188+
cursor_result = await session.execute(cursor_stmt)
189+
cursor = cursor_result.one_or_none()
190+
if cursor is not None:
191+
cursor_index, cursor_id = cursor
192+
if descending:
193+
stmt = stmt.where(
194+
or_(
195+
SessionTurnDBE.turn_index < cursor_index,
196+
and_(
197+
SessionTurnDBE.turn_index == cursor_index,
198+
SessionTurnDBE.id < cursor_id,
199+
),
200+
)
201+
)
202+
else:
203+
stmt = stmt.where(
204+
or_(
205+
SessionTurnDBE.turn_index > cursor_index,
206+
and_(
207+
SessionTurnDBE.turn_index == cursor_index,
208+
SessionTurnDBE.id > cursor_id,
209+
),
210+
)
211+
)
212+
213+
if descending:
214+
stmt = stmt.order_by(
215+
SessionTurnDBE.turn_index.desc(),
216+
SessionTurnDBE.id.desc(),
129217
)
130218
else:
131-
stmt = stmt.order_by(SessionTurnDBE.created_at.desc())
219+
stmt = stmt.order_by(
220+
SessionTurnDBE.turn_index.asc(),
221+
SessionTurnDBE.id.asc(),
222+
)
223+
224+
if windowing and windowing.limit:
225+
stmt = stmt.limit(windowing.limit)
132226

133227
result = await session.execute(stmt)
134228
return [map_turn_dbe_to_dto(turn_dbe=dbe) for dbe in result.scalars().all()]

0 commit comments

Comments
 (0)