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.
- steerable_agent_runtime/__init__.py +34 -0
- steerable_agent_runtime/errors.py +34 -0
- steerable_agent_runtime/llm/__init__.py +104 -0
- steerable_agent_runtime/llm/anthropic_native.py +261 -0
- steerable_agent_runtime/llm/openai_compat.py +256 -0
- steerable_agent_runtime/storage/__init__.py +82 -0
- steerable_agent_runtime/storage/in_memory.py +151 -0
- steerable_agent_runtime/storage/sqlalchemy_store.py +340 -0
- steerable_agent_runtime/tools.py +251 -0
- steerable_agent_runtime/transport/__init__.py +39 -0
- steerable_agent_runtime/transport/fastapi_sse.py +116 -0
- steerable_agent_runtime/transport/stdio_jsonrpc.py +319 -0
- steerable_agent_runtime-0.1.0.dist-info/METADATA +61 -0
- steerable_agent_runtime-0.1.0.dist-info/RECORD +16 -0
- steerable_agent_runtime-0.1.0.dist-info/WHEEL +5 -0
- steerable_agent_runtime-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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
|