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,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)