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,251 @@
|
|
|
1
|
+
"""In-process tool registry.
|
|
2
|
+
|
|
3
|
+
Tools are registered with `@tool` (or `ToolRouter.register`). Each tool gets:
|
|
4
|
+
|
|
5
|
+
* ``name`` — unique identifier (used in ``ToolCall.name``)
|
|
6
|
+
* ``handler`` — callable that takes a dict of arguments and returns a value
|
|
7
|
+
* ``mode`` — `ToolMode` (auto-classified from name unless overridden)
|
|
8
|
+
* ``schema`` — optional JSON Schema for ``arguments`` (surfaced to the LLM)
|
|
9
|
+
* ``description`` — natural-language description for the LLM
|
|
10
|
+
|
|
11
|
+
The router accepts a `ToolCall` and returns a `ToolResult`. Errors raised by
|
|
12
|
+
handlers are wrapped into `ToolResult.error` so the loop can decide whether to
|
|
13
|
+
self-heal.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import asyncio
|
|
19
|
+
import inspect
|
|
20
|
+
import logging
|
|
21
|
+
import time
|
|
22
|
+
from collections.abc import Awaitable, Callable
|
|
23
|
+
from dataclasses import dataclass, field
|
|
24
|
+
from typing import Any
|
|
25
|
+
|
|
26
|
+
from steerable_agent_harness.policy import ToolMode, decide_tool_mode
|
|
27
|
+
from steerable_agent_protocol.generated import ToolCall, ToolResult
|
|
28
|
+
|
|
29
|
+
from .errors import PolicyDeniedError, ToolDispatchError
|
|
30
|
+
|
|
31
|
+
logger = logging.getLogger(__name__)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
ToolHandler = Callable[..., Any] | Callable[..., Awaitable[Any]]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(slots=True)
|
|
38
|
+
class RegisteredTool:
|
|
39
|
+
name: str
|
|
40
|
+
handler: ToolHandler
|
|
41
|
+
mode: ToolMode
|
|
42
|
+
description: str = ""
|
|
43
|
+
schema: dict[str, Any] = field(default_factory=lambda: {"type": "object", "properties": {}})
|
|
44
|
+
require_consent: bool = False
|
|
45
|
+
metadata: dict[str, Any] = field(default_factory=dict)
|
|
46
|
+
|
|
47
|
+
def to_openai_function(self) -> dict[str, Any]:
|
|
48
|
+
return {
|
|
49
|
+
"type": "function",
|
|
50
|
+
"function": {
|
|
51
|
+
"name": self.name,
|
|
52
|
+
"description": self.description,
|
|
53
|
+
"parameters": self.schema,
|
|
54
|
+
},
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
class ToolRouter:
|
|
59
|
+
"""Async-safe in-process tool dispatch."""
|
|
60
|
+
|
|
61
|
+
def __init__(self) -> None:
|
|
62
|
+
self._tools: dict[str, RegisteredTool] = {}
|
|
63
|
+
|
|
64
|
+
# ------------------------------------------------------------------
|
|
65
|
+
# Registration
|
|
66
|
+
# ------------------------------------------------------------------
|
|
67
|
+
|
|
68
|
+
def register(
|
|
69
|
+
self,
|
|
70
|
+
handler: ToolHandler,
|
|
71
|
+
*,
|
|
72
|
+
name: str | None = None,
|
|
73
|
+
mode: ToolMode | None = None,
|
|
74
|
+
description: str | None = None,
|
|
75
|
+
schema: dict[str, Any] | None = None,
|
|
76
|
+
require_consent: bool | None = None,
|
|
77
|
+
metadata: dict[str, Any] | None = None,
|
|
78
|
+
) -> RegisteredTool:
|
|
79
|
+
resolved_name = name or getattr(handler, "__name__", None)
|
|
80
|
+
if not resolved_name:
|
|
81
|
+
raise ToolDispatchError("Tool handler must have a name")
|
|
82
|
+
if resolved_name in self._tools:
|
|
83
|
+
raise ToolDispatchError(f"Tool already registered: {resolved_name}")
|
|
84
|
+
resolved_mode: ToolMode = mode or decide_tool_mode(resolved_name)
|
|
85
|
+
resolved_consent = (
|
|
86
|
+
require_consent if require_consent is not None else resolved_mode == "destructive"
|
|
87
|
+
)
|
|
88
|
+
tool_meta = RegisteredTool(
|
|
89
|
+
name=resolved_name,
|
|
90
|
+
handler=handler,
|
|
91
|
+
mode=resolved_mode,
|
|
92
|
+
description=description or (inspect.getdoc(handler) or "").strip(),
|
|
93
|
+
schema=schema or {"type": "object", "properties": {}},
|
|
94
|
+
require_consent=resolved_consent,
|
|
95
|
+
metadata=dict(metadata or {}),
|
|
96
|
+
)
|
|
97
|
+
self._tools[resolved_name] = tool_meta
|
|
98
|
+
return tool_meta
|
|
99
|
+
|
|
100
|
+
def unregister(self, name: str) -> None:
|
|
101
|
+
self._tools.pop(name, None)
|
|
102
|
+
|
|
103
|
+
def list_tools(self) -> list[RegisteredTool]:
|
|
104
|
+
return list(self._tools.values())
|
|
105
|
+
|
|
106
|
+
def describe(self) -> list[dict[str, Any]]:
|
|
107
|
+
return [t.to_openai_function() for t in self._tools.values()]
|
|
108
|
+
|
|
109
|
+
def get(self, name: str) -> RegisteredTool | None:
|
|
110
|
+
return self._tools.get(name)
|
|
111
|
+
|
|
112
|
+
# ------------------------------------------------------------------
|
|
113
|
+
# Dispatch
|
|
114
|
+
# ------------------------------------------------------------------
|
|
115
|
+
|
|
116
|
+
async def dispatch(
|
|
117
|
+
self,
|
|
118
|
+
call: ToolCall,
|
|
119
|
+
*,
|
|
120
|
+
consent_granted: bool = False,
|
|
121
|
+
context: dict[str, Any] | None = None,
|
|
122
|
+
) -> ToolResult:
|
|
123
|
+
tool = self._tools.get(call.name)
|
|
124
|
+
if tool is None:
|
|
125
|
+
return ToolResult(
|
|
126
|
+
success=False,
|
|
127
|
+
error=f"Unknown tool: {call.name}",
|
|
128
|
+
terminal=False,
|
|
129
|
+
needsFollowup=False,
|
|
130
|
+
)
|
|
131
|
+
if tool.require_consent and not consent_granted:
|
|
132
|
+
raise PolicyDeniedError(
|
|
133
|
+
f"Tool '{tool.name}' requires explicit consent",
|
|
134
|
+
data={"tool": tool.name, "mode": tool.mode},
|
|
135
|
+
)
|
|
136
|
+
started = time.monotonic()
|
|
137
|
+
try:
|
|
138
|
+
result = await self._invoke(tool, call.arguments or {}, context or {})
|
|
139
|
+
except Exception as exc: # noqa: BLE001 — wrap for the loop
|
|
140
|
+
logger.exception("Tool %s failed", tool.name)
|
|
141
|
+
return ToolResult(
|
|
142
|
+
success=False,
|
|
143
|
+
error=str(exc),
|
|
144
|
+
terminal=False,
|
|
145
|
+
needsFollowup=True,
|
|
146
|
+
data={"durationMs": int((time.monotonic() - started) * 1000)},
|
|
147
|
+
)
|
|
148
|
+
return _coerce_to_tool_result(result, duration_ms=int((time.monotonic() - started) * 1000))
|
|
149
|
+
|
|
150
|
+
async def _invoke(
|
|
151
|
+
self,
|
|
152
|
+
tool: RegisteredTool,
|
|
153
|
+
arguments: dict[str, Any],
|
|
154
|
+
context: dict[str, Any],
|
|
155
|
+
) -> Any:
|
|
156
|
+
signature = inspect.signature(tool.handler)
|
|
157
|
+
kwargs: dict[str, Any] = {}
|
|
158
|
+
for parameter in signature.parameters.values():
|
|
159
|
+
if parameter.kind in (
|
|
160
|
+
inspect.Parameter.VAR_POSITIONAL,
|
|
161
|
+
inspect.Parameter.VAR_KEYWORD,
|
|
162
|
+
):
|
|
163
|
+
continue
|
|
164
|
+
name = parameter.name
|
|
165
|
+
if name in arguments:
|
|
166
|
+
kwargs[name] = arguments[name]
|
|
167
|
+
elif name == "context":
|
|
168
|
+
kwargs[name] = context
|
|
169
|
+
result = tool.handler(**kwargs)
|
|
170
|
+
if inspect.isawaitable(result):
|
|
171
|
+
return await result
|
|
172
|
+
if asyncio.iscoroutine(result): # pragma: no cover - defensive
|
|
173
|
+
return await result
|
|
174
|
+
return result
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
# ---------------------------------------------------------------------------
|
|
178
|
+
# Decorator helper
|
|
179
|
+
# ---------------------------------------------------------------------------
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
def tool(
|
|
183
|
+
*,
|
|
184
|
+
name: str | None = None,
|
|
185
|
+
mode: ToolMode | None = None,
|
|
186
|
+
description: str | None = None,
|
|
187
|
+
schema: dict[str, Any] | None = None,
|
|
188
|
+
require_consent: bool | None = None,
|
|
189
|
+
router: ToolRouter | None = None,
|
|
190
|
+
) -> Callable[[ToolHandler], ToolHandler]:
|
|
191
|
+
"""Decorator form of `ToolRouter.register()`.
|
|
192
|
+
|
|
193
|
+
Usage::
|
|
194
|
+
|
|
195
|
+
router = ToolRouter()
|
|
196
|
+
|
|
197
|
+
@tool(router=router, description="List events for the user")
|
|
198
|
+
async def list_events(limit: int = 20) -> list[dict]:
|
|
199
|
+
...
|
|
200
|
+
|
|
201
|
+
If `router` is omitted the decorator stores registration metadata on the
|
|
202
|
+
function as ``__steerable_tool_meta__`` so it can be batch-registered
|
|
203
|
+
later via ``router.register_decorated(handler)``.
|
|
204
|
+
"""
|
|
205
|
+
|
|
206
|
+
def _decorator(handler: ToolHandler) -> ToolHandler:
|
|
207
|
+
meta = {
|
|
208
|
+
"name": name,
|
|
209
|
+
"mode": mode,
|
|
210
|
+
"description": description,
|
|
211
|
+
"schema": schema,
|
|
212
|
+
"require_consent": require_consent,
|
|
213
|
+
}
|
|
214
|
+
setattr(handler, "__steerable_tool_meta__", meta)
|
|
215
|
+
if router is not None:
|
|
216
|
+
router.register(
|
|
217
|
+
handler,
|
|
218
|
+
name=name,
|
|
219
|
+
mode=mode,
|
|
220
|
+
description=description,
|
|
221
|
+
schema=schema,
|
|
222
|
+
require_consent=require_consent,
|
|
223
|
+
)
|
|
224
|
+
return handler
|
|
225
|
+
|
|
226
|
+
return _decorator
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
# ---------------------------------------------------------------------------
|
|
230
|
+
# Result coercion
|
|
231
|
+
# ---------------------------------------------------------------------------
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def _coerce_to_tool_result(value: Any, *, duration_ms: int) -> ToolResult:
|
|
235
|
+
"""Convert handler return value into a `ToolResult`."""
|
|
236
|
+
|
|
237
|
+
if isinstance(value, ToolResult):
|
|
238
|
+
if value.data is None:
|
|
239
|
+
value.data = {}
|
|
240
|
+
value.data.setdefault("durationMs", duration_ms)
|
|
241
|
+
return value
|
|
242
|
+
if isinstance(value, dict) and "success" in value:
|
|
243
|
+
result = ToolResult(**value)
|
|
244
|
+
if result.data is None:
|
|
245
|
+
result.data = {}
|
|
246
|
+
result.data.setdefault("durationMs", duration_ms)
|
|
247
|
+
return result
|
|
248
|
+
return ToolResult(
|
|
249
|
+
success=True,
|
|
250
|
+
data={"value": value, "durationMs": duration_ms},
|
|
251
|
+
)
|
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
"""TransportAdapter interface + reference implementations."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import AsyncIterator
|
|
6
|
+
from typing import Any, Protocol, runtime_checkable
|
|
7
|
+
|
|
8
|
+
from steerable_agent_protocol.generated import SSEEvent
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@runtime_checkable
|
|
12
|
+
class TransportAdapter(Protocol):
|
|
13
|
+
"""Abstract bidirectional transport.
|
|
14
|
+
|
|
15
|
+
The runtime emits `SSEEvent` instances via `emit()` and receives requests
|
|
16
|
+
via the implementation-specific entrypoint (HTTP route handler, stdio
|
|
17
|
+
pump, websocket message, etc.).
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
async def emit(self, event: SSEEvent) -> None: ...
|
|
21
|
+
|
|
22
|
+
async def aclose(self) -> None: ...
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
from .fastapi_sse import FastAPISseTransport, sse_response # noqa: E402
|
|
26
|
+
from .stdio_jsonrpc import ( # noqa: E402
|
|
27
|
+
StdioJsonRpcTransport,
|
|
28
|
+
JsonRpcMethodHandler,
|
|
29
|
+
JsonRpcServer,
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
__all__ = [
|
|
33
|
+
"TransportAdapter",
|
|
34
|
+
"FastAPISseTransport",
|
|
35
|
+
"sse_response",
|
|
36
|
+
"StdioJsonRpcTransport",
|
|
37
|
+
"JsonRpcMethodHandler",
|
|
38
|
+
"JsonRpcServer",
|
|
39
|
+
]
|
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
"""Server-Sent Events transport built on top of FastAPI / Starlette."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import json
|
|
7
|
+
from collections.abc import AsyncIterator
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
from steerable_agent_protocol.generated import SSEEvent
|
|
11
|
+
|
|
12
|
+
from ..errors import TransportError
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class FastAPISseTransport:
|
|
16
|
+
"""Per-request SSE transport.
|
|
17
|
+
|
|
18
|
+
Usage::
|
|
19
|
+
|
|
20
|
+
@app.post("/chat/run")
|
|
21
|
+
async def run_chat(...):
|
|
22
|
+
transport = FastAPISseTransport()
|
|
23
|
+
asyncio.create_task(_run_loop(transport, ...))
|
|
24
|
+
return await sse_response(transport)
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(self, *, queue_size: int = 256) -> None:
|
|
28
|
+
self._queue: asyncio.Queue[SSEEvent | None] = asyncio.Queue(maxsize=queue_size)
|
|
29
|
+
self._closed = False
|
|
30
|
+
|
|
31
|
+
async def emit(self, event: SSEEvent) -> None:
|
|
32
|
+
if self._closed:
|
|
33
|
+
raise TransportError("Transport is already closed")
|
|
34
|
+
await self._queue.put(event)
|
|
35
|
+
|
|
36
|
+
async def aclose(self) -> None:
|
|
37
|
+
if self._closed:
|
|
38
|
+
return
|
|
39
|
+
self._closed = True
|
|
40
|
+
await self._queue.put(None)
|
|
41
|
+
|
|
42
|
+
async def stream(self) -> AsyncIterator[SSEEvent]:
|
|
43
|
+
while True:
|
|
44
|
+
event = await self._queue.get()
|
|
45
|
+
if event is None:
|
|
46
|
+
return
|
|
47
|
+
yield event
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
async def sse_response(
|
|
51
|
+
transport: FastAPISseTransport,
|
|
52
|
+
*,
|
|
53
|
+
media_type: str = "text/event-stream",
|
|
54
|
+
headers: dict[str, str] | None = None,
|
|
55
|
+
):
|
|
56
|
+
"""Build a Starlette ``StreamingResponse`` that drains ``transport``.
|
|
57
|
+
|
|
58
|
+
The response will keep the connection open until ``transport.aclose()`` is
|
|
59
|
+
called from the producing task.
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
try:
|
|
63
|
+
from starlette.responses import StreamingResponse # local import (optional dep)
|
|
64
|
+
except Exception as exc: # pragma: no cover
|
|
65
|
+
raise ImportError(
|
|
66
|
+
"FastAPISseTransport.stream requires `starlette` (install "
|
|
67
|
+
"`steerable-agent-runtime[fastapi]`)"
|
|
68
|
+
) from exc
|
|
69
|
+
|
|
70
|
+
async def _gen() -> AsyncIterator[bytes]:
|
|
71
|
+
async for event in transport.stream():
|
|
72
|
+
payload = encode_sse_event(event)
|
|
73
|
+
yield payload.encode("utf-8")
|
|
74
|
+
|
|
75
|
+
response_headers = {
|
|
76
|
+
"Cache-Control": "no-cache",
|
|
77
|
+
"Connection": "keep-alive",
|
|
78
|
+
"X-Accel-Buffering": "no",
|
|
79
|
+
}
|
|
80
|
+
if headers:
|
|
81
|
+
response_headers.update(headers)
|
|
82
|
+
return StreamingResponse(_gen(), media_type=media_type, headers=response_headers)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def encode_sse_event(event: SSEEvent) -> str:
|
|
86
|
+
"""Serialize an `SSEEvent` to a wire-format `data: <json>\n\n` block."""
|
|
87
|
+
|
|
88
|
+
body = event.model_dump(exclude_none=True)
|
|
89
|
+
payload = json.dumps(body, ensure_ascii=False, separators=(",", ":"))
|
|
90
|
+
lines = []
|
|
91
|
+
if event.event:
|
|
92
|
+
lines.append(f"event: {event.event}")
|
|
93
|
+
lines.append(f"data: {payload}")
|
|
94
|
+
return "\n".join(lines) + "\n\n"
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def decode_sse_event(raw: str) -> SSEEvent | None:
|
|
98
|
+
"""Parse a `data: <json>` block back into `SSEEvent`. Returns None when the
|
|
99
|
+
block is empty or malformed (mirrors browser EventSource behavior)."""
|
|
100
|
+
|
|
101
|
+
data: str | None = None
|
|
102
|
+
event: str | None = None
|
|
103
|
+
for line in raw.splitlines():
|
|
104
|
+
if line.startswith("data:"):
|
|
105
|
+
data = (data or "") + line[5:].lstrip()
|
|
106
|
+
elif line.startswith("event:"):
|
|
107
|
+
event = line[6:].strip()
|
|
108
|
+
if data is None:
|
|
109
|
+
return None
|
|
110
|
+
try:
|
|
111
|
+
payload = json.loads(data)
|
|
112
|
+
except json.JSONDecodeError:
|
|
113
|
+
return None
|
|
114
|
+
if event and "event" not in payload:
|
|
115
|
+
payload["event"] = event
|
|
116
|
+
return SSEEvent.model_validate(payload)
|