steerable-agent-runtime 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,82 @@
1
+ """StorageAdapter interface + reference implementations."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Iterable, Sequence
6
+ from typing import Any, Protocol, runtime_checkable
7
+
8
+ from steerable_agent_protocol.generated import (
9
+ AgentSession,
10
+ ChatAgent,
11
+ ChatMessage,
12
+ HarnessTrace,
13
+ TraceEvent,
14
+ TraceSpan,
15
+ )
16
+
17
+
18
+ @runtime_checkable
19
+ class StorageAdapter(Protocol):
20
+ """Persistence interface for the runtime.
21
+
22
+ Implementations must be **safe under concurrent ``await``** but are not
23
+ required to be process-safe. The reference SQLAlchemy adapter delegates
24
+ isolation to the underlying database.
25
+ """
26
+
27
+ # -- AgentSession ---------------------------------------------------
28
+
29
+ async def upsert_session(self, session: AgentSession) -> AgentSession: ...
30
+
31
+ async def get_session(self, session_id: str) -> AgentSession | None: ...
32
+
33
+ async def list_sessions(
34
+ self,
35
+ *,
36
+ user_id: str | None = None,
37
+ chat_id: str | None = None,
38
+ active_only: bool = False,
39
+ ) -> list[AgentSession]: ...
40
+
41
+ # -- ChatAgent ------------------------------------------------------
42
+
43
+ async def upsert_agent(self, agent: ChatAgent) -> ChatAgent: ...
44
+
45
+ async def get_agent(self, agent_id: str) -> ChatAgent | None: ...
46
+
47
+ async def list_agents(self, *, include_archived: bool = False) -> list[ChatAgent]: ...
48
+
49
+ # -- ChatMessage ----------------------------------------------------
50
+
51
+ async def append_message(self, message: ChatMessage) -> ChatMessage: ...
52
+
53
+ async def list_messages(self, chat_id: str, *, limit: int | None = None) -> list[ChatMessage]: ...
54
+
55
+ # -- HarnessTrace + spans + events ---------------------------------
56
+
57
+ async def upsert_trace(self, trace: HarnessTrace) -> HarnessTrace: ...
58
+
59
+ async def get_trace(self, trace_id: str) -> HarnessTrace | None: ...
60
+
61
+ async def append_spans(self, trace_id: str, spans: Iterable[TraceSpan]) -> None: ...
62
+
63
+ async def list_spans(self, trace_id: str) -> list[TraceSpan]: ...
64
+
65
+ async def append_events(self, trace_id: str, events: Iterable[TraceEvent]) -> None: ...
66
+
67
+ async def list_events(self, trace_id: str) -> list[TraceEvent]: ...
68
+
69
+
70
+ from .in_memory import InMemoryStorage # noqa: E402
71
+
72
+ try:
73
+ from .sqlalchemy_store import SqlAlchemyStorage # noqa: F401
74
+ except Exception: # pragma: no cover - optional dep
75
+ SqlAlchemyStorage = None # type: ignore[assignment]
76
+
77
+
78
+ __all__ = [
79
+ "StorageAdapter",
80
+ "InMemoryStorage",
81
+ "SqlAlchemyStorage",
82
+ ]
@@ -0,0 +1,151 @@
1
+ """Reference in-memory StorageAdapter (default for sidecar / dev)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ from collections.abc import Iterable
7
+ from copy import deepcopy
8
+
9
+ from steerable_agent_protocol.generated import (
10
+ AgentSession,
11
+ ChatAgent,
12
+ ChatMessage,
13
+ HarnessTrace,
14
+ TraceEvent,
15
+ TraceSpan,
16
+ )
17
+
18
+ from ..errors import StorageError
19
+
20
+
21
+ class InMemoryStorage:
22
+ """Thread-safe in-memory storage. All mutations happen under an asyncio
23
+ lock so concurrent dispatch from a single event loop is safe."""
24
+
25
+ def __init__(self) -> None:
26
+ self._lock = asyncio.Lock()
27
+ self._sessions: dict[str, AgentSession] = {}
28
+ self._agents: dict[str, ChatAgent] = {}
29
+ self._messages: dict[str, list[ChatMessage]] = {}
30
+ self._traces: dict[str, HarnessTrace] = {}
31
+ self._spans: dict[str, list[TraceSpan]] = {}
32
+ self._events: dict[str, list[TraceEvent]] = {}
33
+
34
+ # ------------------------------------------------------------------
35
+ # Sessions
36
+ # ------------------------------------------------------------------
37
+
38
+ async def upsert_session(self, session: AgentSession) -> AgentSession:
39
+ async with self._lock:
40
+ self._sessions[session.sessionId] = deepcopy(session)
41
+ return deepcopy(session)
42
+
43
+ async def get_session(self, session_id: str) -> AgentSession | None:
44
+ async with self._lock:
45
+ value = self._sessions.get(session_id)
46
+ return deepcopy(value) if value else None
47
+
48
+ async def list_sessions(
49
+ self,
50
+ *,
51
+ user_id: str | None = None,
52
+ chat_id: str | None = None,
53
+ active_only: bool = False,
54
+ ) -> list[AgentSession]:
55
+ async with self._lock:
56
+ sessions = list(self._sessions.values())
57
+ if user_id is not None:
58
+ sessions = [s for s in sessions if s.userId == user_id]
59
+ if chat_id is not None:
60
+ sessions = [s for s in sessions if s.chatId == chat_id]
61
+ if active_only:
62
+ sessions = [s for s in sessions if s.isActive]
63
+ sessions.sort(key=lambda s: s.updatedAt, reverse=True)
64
+ return [deepcopy(s) for s in sessions]
65
+
66
+ # ------------------------------------------------------------------
67
+ # Agents
68
+ # ------------------------------------------------------------------
69
+
70
+ async def upsert_agent(self, agent: ChatAgent) -> ChatAgent:
71
+ async with self._lock:
72
+ self._agents[agent.id] = deepcopy(agent)
73
+ return deepcopy(agent)
74
+
75
+ async def get_agent(self, agent_id: str) -> ChatAgent | None:
76
+ async with self._lock:
77
+ value = self._agents.get(agent_id)
78
+ return deepcopy(value) if value else None
79
+
80
+ async def list_agents(self, *, include_archived: bool = False) -> list[ChatAgent]:
81
+ async with self._lock:
82
+ agents = list(self._agents.values())
83
+ if not include_archived:
84
+ agents = [a for a in agents if not a.isArchived]
85
+ agents.sort(key=lambda a: (a.sortOrder, a.createdAt))
86
+ return [deepcopy(a) for a in agents]
87
+
88
+ # ------------------------------------------------------------------
89
+ # Messages
90
+ # ------------------------------------------------------------------
91
+
92
+ async def append_message(self, message: ChatMessage) -> ChatMessage:
93
+ if not message.chatId:
94
+ raise StorageError("ChatMessage.chatId is required for append_message")
95
+ async with self._lock:
96
+ bucket = self._messages.setdefault(message.chatId, [])
97
+ bucket.append(deepcopy(message))
98
+ return deepcopy(message)
99
+
100
+ async def list_messages(
101
+ self, chat_id: str, *, limit: int | None = None
102
+ ) -> list[ChatMessage]:
103
+ async with self._lock:
104
+ bucket = list(self._messages.get(chat_id, []))
105
+ bucket.sort(key=lambda m: m.createdAt)
106
+ if limit is not None:
107
+ bucket = bucket[-limit:]
108
+ return [deepcopy(m) for m in bucket]
109
+
110
+ # ------------------------------------------------------------------
111
+ # Traces / spans / events
112
+ # ------------------------------------------------------------------
113
+
114
+ async def upsert_trace(self, trace: HarnessTrace) -> HarnessTrace:
115
+ async with self._lock:
116
+ self._traces[trace.traceId] = deepcopy(trace)
117
+ return deepcopy(trace)
118
+
119
+ async def get_trace(self, trace_id: str) -> HarnessTrace | None:
120
+ async with self._lock:
121
+ value = self._traces.get(trace_id)
122
+ return deepcopy(value) if value else None
123
+
124
+ async def append_spans(self, trace_id: str, spans: Iterable[TraceSpan]) -> None:
125
+ async with self._lock:
126
+ bucket = self._spans.setdefault(trace_id, [])
127
+ for span in spans:
128
+ bucket.append(deepcopy(span))
129
+ trace = self._traces.get(trace_id)
130
+ if trace is not None:
131
+ trace.spanCount = len(bucket)
132
+
133
+ async def list_spans(self, trace_id: str) -> list[TraceSpan]:
134
+ async with self._lock:
135
+ return [deepcopy(span) for span in self._spans.get(trace_id, [])]
136
+
137
+ async def append_events(self, trace_id: str, events: Iterable[TraceEvent]) -> None:
138
+ async with self._lock:
139
+ bucket = self._events.setdefault(trace_id, [])
140
+ for event in events:
141
+ bucket.append(deepcopy(event))
142
+ trace = self._traces.get(trace_id)
143
+ if trace is not None:
144
+ trace.eventCount = len(bucket)
145
+
146
+ async def list_events(self, trace_id: str) -> list[TraceEvent]:
147
+ async with self._lock:
148
+ return sorted(
149
+ [deepcopy(event) for event in self._events.get(trace_id, [])],
150
+ key=lambda event: event.sequence,
151
+ )
@@ -0,0 +1,340 @@
1
+ """SQLAlchemy-backed StorageAdapter (optional)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Iterable
6
+ from typing import Any
7
+
8
+ from steerable_agent_protocol.generated import (
9
+ AgentSession,
10
+ ChatAgent,
11
+ ChatMessage,
12
+ HarnessTrace,
13
+ TraceEvent,
14
+ TraceSpan,
15
+ )
16
+
17
+ try:
18
+ from sqlalchemy import (
19
+ JSON,
20
+ Boolean,
21
+ Column,
22
+ DateTime,
23
+ Integer,
24
+ MetaData,
25
+ String,
26
+ Table,
27
+ Text,
28
+ UniqueConstraint,
29
+ select,
30
+ )
31
+ from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker
32
+ except Exception as exc: # pragma: no cover - optional dep
33
+ raise ImportError(
34
+ "SqlAlchemyStorage requires sqlalchemy>=2.0. Install with "
35
+ "`pip install steerable-agent-runtime[sqlalchemy]`."
36
+ ) from exc
37
+
38
+ from ..errors import StorageError
39
+
40
+ metadata = MetaData()
41
+
42
+
43
+ sessions_table = Table(
44
+ "steerable_session",
45
+ metadata,
46
+ Column("sessionId", String(191), primary_key=True),
47
+ Column("userId", String(191), nullable=False, index=True),
48
+ Column("projectId", String(191), nullable=True),
49
+ Column("chatId", String(191), nullable=False, index=True),
50
+ Column("currentStage", String(191), nullable=False),
51
+ Column("nextStage", String(191), nullable=True),
52
+ Column("scenario", String(191), nullable=False, default="agent-entry"),
53
+ Column("stageData", JSON, nullable=True),
54
+ Column("isActive", Boolean, nullable=False, default=True),
55
+ Column("createdAt", DateTime, nullable=False),
56
+ Column("updatedAt", DateTime, nullable=False),
57
+ )
58
+
59
+ agents_table = Table(
60
+ "steerable_agent",
61
+ metadata,
62
+ Column("id", String(191), primary_key=True),
63
+ Column("slug", String(191), nullable=True),
64
+ Column("name", String(191), nullable=False),
65
+ Column("icon", String(191), nullable=True),
66
+ Column("color", String(191), nullable=True),
67
+ Column("description", Text, nullable=True),
68
+ Column("rolePrompt", Text, nullable=True),
69
+ Column("forbiddenPrompt", Text, nullable=True),
70
+ Column("skillIds", JSON, nullable=False, default=list),
71
+ Column("allowExternalSkills", Boolean, nullable=False, default=True),
72
+ Column("isBuiltin", Boolean, nullable=False, default=False),
73
+ Column("isArchived", Boolean, nullable=False, default=False),
74
+ Column("sortOrder", Integer, nullable=False, default=0),
75
+ Column("createdAt", DateTime, nullable=False),
76
+ Column("updatedAt", DateTime, nullable=False),
77
+ )
78
+
79
+ messages_table = Table(
80
+ "steerable_message",
81
+ metadata,
82
+ Column("id", String(191), primary_key=True),
83
+ Column("chatId", String(191), nullable=False, index=True),
84
+ Column("role", String(32), nullable=False),
85
+ Column("content", Text, nullable=False),
86
+ Column("agentId", String(191), nullable=True),
87
+ Column("toolCalls", JSON, nullable=True),
88
+ Column("toolResult", JSON, nullable=True),
89
+ Column("createdAt", DateTime, nullable=False, index=True),
90
+ Column("updatedAt", DateTime, nullable=True),
91
+ )
92
+
93
+ traces_table = Table(
94
+ "steerable_trace",
95
+ metadata,
96
+ Column("traceId", String(191), primary_key=True),
97
+ Column("userId", String(191), nullable=True, index=True),
98
+ Column("chatId", String(191), nullable=True, index=True),
99
+ Column("sessionId", String(191), nullable=True, index=True),
100
+ Column("assistantMessageId", String(191), nullable=True),
101
+ Column("status", String(32), nullable=False, default="running"),
102
+ Column("durationMs", Integer, nullable=True),
103
+ Column("hadError", Boolean, nullable=False, default=False),
104
+ Column("errorMessage", String(2048), nullable=True),
105
+ Column("eventCount", Integer, nullable=False, default=0),
106
+ Column("spanCount", Integer, nullable=False, default=0),
107
+ Column("totalTokens", Integer, nullable=True),
108
+ Column("modelId", String(191), nullable=True),
109
+ Column("startedAtMs", Integer, nullable=True),
110
+ Column("createdAt", DateTime, nullable=False),
111
+ Column("updatedAt", DateTime, nullable=False),
112
+ )
113
+
114
+ spans_table = Table(
115
+ "steerable_trace_span",
116
+ metadata,
117
+ Column("spanId", String(191), primary_key=True),
118
+ Column("traceId", String(191), nullable=False, index=True),
119
+ Column("parentSpanId", String(191), nullable=True),
120
+ Column("name", String(191), nullable=False),
121
+ Column("kind", String(32), nullable=False, default="custom"),
122
+ Column("startMs", Integer, nullable=False),
123
+ Column("endMs", Integer, nullable=True),
124
+ Column("durationMs", Integer, nullable=True),
125
+ Column("status", String(32), nullable=False, default="running"),
126
+ Column("attrs", JSON, nullable=False, default=dict),
127
+ )
128
+
129
+ events_table = Table(
130
+ "steerable_trace_event",
131
+ metadata,
132
+ Column("id", String(191), primary_key=True),
133
+ Column("traceId", String(191), nullable=False, index=True),
134
+ Column("kind", String(32), nullable=False),
135
+ Column("name", String(191), nullable=False),
136
+ Column("sequence", Integer, nullable=False),
137
+ Column("timestampMs", Integer, nullable=False),
138
+ Column("durationMs", Integer, nullable=True),
139
+ Column("status", String(32), nullable=True),
140
+ Column("payload", JSON, nullable=True),
141
+ Column("createdAt", DateTime, nullable=True),
142
+ UniqueConstraint("traceId", "sequence", name="steerable_trace_event_seq_uniq"),
143
+ )
144
+
145
+
146
+ def _model_to_row(model: Any) -> dict[str, Any]:
147
+ return model.model_dump()
148
+
149
+
150
+ class SqlAlchemyStorage:
151
+ """SQLAlchemy 2.0 async StorageAdapter."""
152
+
153
+ def __init__(self, engine: AsyncEngine) -> None:
154
+ self._engine = engine
155
+ self._sessionmaker = async_sessionmaker(engine, expire_on_commit=False)
156
+
157
+ async def create_all(self) -> None:
158
+ async with self._engine.begin() as conn:
159
+ await conn.run_sync(metadata.create_all)
160
+
161
+ # ------------------------------------------------------------------
162
+ # Sessions
163
+ # ------------------------------------------------------------------
164
+
165
+ async def upsert_session(self, session: AgentSession) -> AgentSession:
166
+ await self._upsert(sessions_table, "sessionId", session.sessionId, _model_to_row(session))
167
+ return session
168
+
169
+ async def get_session(self, session_id: str) -> AgentSession | None:
170
+ row = await self._get_one(sessions_table, sessions_table.c.sessionId == session_id)
171
+ return AgentSession(**row) if row else None
172
+
173
+ async def list_sessions(
174
+ self,
175
+ *,
176
+ user_id: str | None = None,
177
+ chat_id: str | None = None,
178
+ active_only: bool = False,
179
+ ) -> list[AgentSession]:
180
+ clauses = []
181
+ if user_id is not None:
182
+ clauses.append(sessions_table.c.userId == user_id)
183
+ if chat_id is not None:
184
+ clauses.append(sessions_table.c.chatId == chat_id)
185
+ if active_only:
186
+ clauses.append(sessions_table.c.isActive.is_(True))
187
+ rows = await self._select_many(
188
+ sessions_table,
189
+ clauses=clauses,
190
+ order_by=[sessions_table.c.updatedAt.desc()],
191
+ )
192
+ return [AgentSession(**row) for row in rows]
193
+
194
+ # ------------------------------------------------------------------
195
+ # Agents
196
+ # ------------------------------------------------------------------
197
+
198
+ async def upsert_agent(self, agent: ChatAgent) -> ChatAgent:
199
+ await self._upsert(agents_table, "id", agent.id, _model_to_row(agent))
200
+ return agent
201
+
202
+ async def get_agent(self, agent_id: str) -> ChatAgent | None:
203
+ row = await self._get_one(agents_table, agents_table.c.id == agent_id)
204
+ return ChatAgent(**row) if row else None
205
+
206
+ async def list_agents(self, *, include_archived: bool = False) -> list[ChatAgent]:
207
+ clauses = []
208
+ if not include_archived:
209
+ clauses.append(agents_table.c.isArchived.is_(False))
210
+ rows = await self._select_many(
211
+ agents_table,
212
+ clauses=clauses,
213
+ order_by=[agents_table.c.sortOrder.asc(), agents_table.c.createdAt.asc()],
214
+ )
215
+ return [ChatAgent(**row) for row in rows]
216
+
217
+ # ------------------------------------------------------------------
218
+ # Messages
219
+ # ------------------------------------------------------------------
220
+
221
+ async def append_message(self, message: ChatMessage) -> ChatMessage:
222
+ if not message.chatId:
223
+ raise StorageError("ChatMessage.chatId is required")
224
+ await self._insert(messages_table, _model_to_row(message))
225
+ return message
226
+
227
+ async def list_messages(self, chat_id: str, *, limit: int | None = None) -> list[ChatMessage]:
228
+ rows = await self._select_many(
229
+ messages_table,
230
+ clauses=[messages_table.c.chatId == chat_id],
231
+ order_by=[messages_table.c.createdAt.asc()],
232
+ limit=limit,
233
+ )
234
+ return [ChatMessage(**row) for row in rows]
235
+
236
+ # ------------------------------------------------------------------
237
+ # Traces / spans / events
238
+ # ------------------------------------------------------------------
239
+
240
+ async def upsert_trace(self, trace: HarnessTrace) -> HarnessTrace:
241
+ await self._upsert(traces_table, "traceId", trace.traceId, _model_to_row(trace))
242
+ return trace
243
+
244
+ async def get_trace(self, trace_id: str) -> HarnessTrace | None:
245
+ row = await self._get_one(traces_table, traces_table.c.traceId == trace_id)
246
+ return HarnessTrace(**row) if row else None
247
+
248
+ async def append_spans(self, trace_id: str, spans: Iterable[TraceSpan]) -> None:
249
+ rows = []
250
+ for span in spans:
251
+ data = _model_to_row(span)
252
+ data["traceId"] = trace_id
253
+ rows.append(data)
254
+ if rows:
255
+ await self._insert_many(spans_table, rows)
256
+
257
+ async def list_spans(self, trace_id: str) -> list[TraceSpan]:
258
+ rows = await self._select_many(
259
+ spans_table,
260
+ clauses=[spans_table.c.traceId == trace_id],
261
+ order_by=[spans_table.c.startMs.asc()],
262
+ )
263
+ return [TraceSpan(**row) for row in rows]
264
+
265
+ async def append_events(self, trace_id: str, events: Iterable[TraceEvent]) -> None:
266
+ rows = []
267
+ for event in events:
268
+ data = _model_to_row(event)
269
+ data["traceId"] = trace_id
270
+ rows.append(data)
271
+ if rows:
272
+ await self._insert_many(events_table, rows)
273
+
274
+ async def list_events(self, trace_id: str) -> list[TraceEvent]:
275
+ rows = await self._select_many(
276
+ events_table,
277
+ clauses=[events_table.c.traceId == trace_id],
278
+ order_by=[events_table.c.sequence.asc()],
279
+ )
280
+ return [TraceEvent(**row) for row in rows]
281
+
282
+ # ------------------------------------------------------------------
283
+ # Internal helpers
284
+ # ------------------------------------------------------------------
285
+
286
+ async def _upsert(
287
+ self,
288
+ table: Table,
289
+ key_column: str,
290
+ key_value: Any,
291
+ row: dict[str, Any],
292
+ ) -> None:
293
+ async with self._sessionmaker() as session:
294
+ existing = await session.execute(
295
+ select(table).where(getattr(table.c, key_column) == key_value)
296
+ )
297
+ if existing.first():
298
+ await session.execute(
299
+ table.update()
300
+ .where(getattr(table.c, key_column) == key_value)
301
+ .values(**row)
302
+ )
303
+ else:
304
+ await session.execute(table.insert().values(**row))
305
+ await session.commit()
306
+
307
+ async def _insert(self, table: Table, row: dict[str, Any]) -> None:
308
+ async with self._sessionmaker() as session:
309
+ await session.execute(table.insert().values(**row))
310
+ await session.commit()
311
+
312
+ async def _insert_many(self, table: Table, rows: list[dict[str, Any]]) -> None:
313
+ async with self._sessionmaker() as session:
314
+ await session.execute(table.insert(), rows)
315
+ await session.commit()
316
+
317
+ async def _select_many(
318
+ self,
319
+ table: Table,
320
+ *,
321
+ clauses: list[Any] | None = None,
322
+ order_by: list[Any] | None = None,
323
+ limit: int | None = None,
324
+ ) -> list[dict[str, Any]]:
325
+ stmt = select(table)
326
+ for clause in clauses or []:
327
+ stmt = stmt.where(clause)
328
+ for ordering in order_by or []:
329
+ stmt = stmt.order_by(ordering)
330
+ if limit is not None:
331
+ stmt = stmt.limit(limit)
332
+ async with self._sessionmaker() as session:
333
+ result = await session.execute(stmt)
334
+ return [dict(row._mapping) for row in result.fetchall()]
335
+
336
+ async def _get_one(self, table: Table, clause: Any) -> dict[str, Any] | None:
337
+ async with self._sessionmaker() as session:
338
+ result = await session.execute(select(table).where(clause))
339
+ row = result.first()
340
+ return dict(row._mapping) if row else None