assistant-runtime-sdk 1.0.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.
- assistant_runtime_sdk/__init__.py +235 -0
- assistant_runtime_sdk/async_client.py +1989 -0
- assistant_runtime_sdk/auth.py +161 -0
- assistant_runtime_sdk/base.py +2277 -0
- assistant_runtime_sdk/client.py +3039 -0
- assistant_runtime_sdk/exceptions.py +146 -0
- assistant_runtime_sdk/skills.py +328 -0
- assistant_runtime_sdk/streaming.py +223 -0
- assistant_runtime_sdk/types.py +546 -0
- assistant_runtime_sdk-1.0.0.dist-info/METADATA +259 -0
- assistant_runtime_sdk-1.0.0.dist-info/RECORD +13 -0
- assistant_runtime_sdk-1.0.0.dist-info/WHEEL +4 -0
- assistant_runtime_sdk-1.0.0.dist-info/licenses/LICENSE +17 -0
|
@@ -0,0 +1,146 @@
|
|
|
1
|
+
# Assistant Runtime SDK - Exceptions
|
|
2
|
+
# Copyright (C) 2025 Paul Clinton
|
|
3
|
+
# AGPL-3.0 License
|
|
4
|
+
|
|
5
|
+
"""
|
|
6
|
+
Custom exceptions for the Assistant Runtime SDK.
|
|
7
|
+
|
|
8
|
+
All exceptions inherit from ARError for easy catching.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from typing import Optional, List
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class ARError(Exception):
|
|
15
|
+
"""Base exception for all Assistant Runtime SDK errors."""
|
|
16
|
+
|
|
17
|
+
def __init__(self, message: str, error_code: Optional[str] = None):
|
|
18
|
+
super().__init__(message)
|
|
19
|
+
self.message = message
|
|
20
|
+
self.error_code = error_code
|
|
21
|
+
|
|
22
|
+
def __str__(self) -> str:
|
|
23
|
+
if self.error_code:
|
|
24
|
+
return f"[{self.error_code}] {self.message}"
|
|
25
|
+
return self.message
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class ARAuthenticationError(ARError):
|
|
29
|
+
"""
|
|
30
|
+
HMAC signature validation failed.
|
|
31
|
+
|
|
32
|
+
This typically indicates:
|
|
33
|
+
- Invalid tenant_secret
|
|
34
|
+
- Clock skew between client and server
|
|
35
|
+
- Tampered request parameters
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
def __init__(self, message: str = "Authentication failed"):
|
|
39
|
+
super().__init__(message, error_code="AUTH_FAILED")
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class ARRateLimitError(ARError):
|
|
43
|
+
"""
|
|
44
|
+
Rate limit exceeded.
|
|
45
|
+
|
|
46
|
+
Contains retry_after indicating seconds until retry is possible.
|
|
47
|
+
May include models_checked if using auto-model selection.
|
|
48
|
+
"""
|
|
49
|
+
|
|
50
|
+
def __init__(
|
|
51
|
+
self,
|
|
52
|
+
message: str = "Rate limit exceeded",
|
|
53
|
+
retry_after: Optional[float] = None,
|
|
54
|
+
models_checked: Optional[List[str]] = None,
|
|
55
|
+
):
|
|
56
|
+
super().__init__(message, error_code="RATE_LIMITED")
|
|
57
|
+
self.retry_after = retry_after
|
|
58
|
+
self.models_checked = models_checked or []
|
|
59
|
+
|
|
60
|
+
def __str__(self) -> str:
|
|
61
|
+
base = super().__str__()
|
|
62
|
+
if self.retry_after:
|
|
63
|
+
return f"{base} (retry after {self.retry_after:.1f}s)"
|
|
64
|
+
return base
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class ARStreamError(ARError):
|
|
68
|
+
"""
|
|
69
|
+
SSE streaming error.
|
|
70
|
+
|
|
71
|
+
Raised when:
|
|
72
|
+
- Connection drops during streaming
|
|
73
|
+
- Malformed SSE data received
|
|
74
|
+
- Server sends error event
|
|
75
|
+
"""
|
|
76
|
+
|
|
77
|
+
def __init__(self, message: str = "Streaming error", error_code: Optional[str] = None):
|
|
78
|
+
super().__init__(message, error_code=error_code or "STREAM_ERROR")
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
class ARConfigurationError(ARError):
|
|
82
|
+
"""
|
|
83
|
+
Invalid configuration.
|
|
84
|
+
|
|
85
|
+
Raised when:
|
|
86
|
+
- Missing required parameters (tenant_id, tenant_secret)
|
|
87
|
+
- Invalid AR URL format
|
|
88
|
+
- Missing user_id where required
|
|
89
|
+
"""
|
|
90
|
+
|
|
91
|
+
def __init__(self, message: str = "Configuration error"):
|
|
92
|
+
super().__init__(message, error_code="CONFIG_ERROR")
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class ARAPIError(ARError):
|
|
96
|
+
"""
|
|
97
|
+
API request failed.
|
|
98
|
+
|
|
99
|
+
Contains HTTP status code and response details.
|
|
100
|
+
"""
|
|
101
|
+
|
|
102
|
+
def __init__(
|
|
103
|
+
self,
|
|
104
|
+
message: str,
|
|
105
|
+
status_code: Optional[int] = None,
|
|
106
|
+
response_data: Optional[dict] = None,
|
|
107
|
+
):
|
|
108
|
+
error_code = f"HTTP_{status_code}" if status_code else "API_ERROR"
|
|
109
|
+
super().__init__(message, error_code=error_code)
|
|
110
|
+
self.status_code = status_code
|
|
111
|
+
self.response_data = response_data or {}
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class ARTimeoutError(ARError):
|
|
115
|
+
"""
|
|
116
|
+
Request timed out.
|
|
117
|
+
|
|
118
|
+
Raised when connection or read timeout is exceeded.
|
|
119
|
+
"""
|
|
120
|
+
|
|
121
|
+
def __init__(self, message: str = "Request timed out"):
|
|
122
|
+
super().__init__(message, error_code="TIMEOUT")
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class ARConnectionError(ARError):
|
|
126
|
+
"""
|
|
127
|
+
Connection failed.
|
|
128
|
+
|
|
129
|
+
Raised when unable to connect to Assistant Runtime server.
|
|
130
|
+
"""
|
|
131
|
+
|
|
132
|
+
def __init__(self, message: str = "Connection failed"):
|
|
133
|
+
super().__init__(message, error_code="CONNECTION_ERROR")
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
class ARBillingUnavailableError(ARError):
|
|
137
|
+
"""
|
|
138
|
+
Billing features are not available on this server.
|
|
139
|
+
|
|
140
|
+
Raised when:
|
|
141
|
+
- The payments app is not installed on the target server
|
|
142
|
+
- check_billing_available() returned False and a billing method was called
|
|
143
|
+
"""
|
|
144
|
+
|
|
145
|
+
def __init__(self, message: str = "Billing is not available"):
|
|
146
|
+
super().__init__(message, error_code="BILLING_UNAVAILABLE")
|
|
@@ -0,0 +1,328 @@
|
|
|
1
|
+
# Assistant Runtime SDK - Skill Provider
|
|
2
|
+
# Copyright (C) 2025 Paul Clinton
|
|
3
|
+
# AGPL-3.0 License
|
|
4
|
+
|
|
5
|
+
"""
|
|
6
|
+
Skill Provider - Bridges MCP resources to Strands Agent.
|
|
7
|
+
|
|
8
|
+
When MCP server has resources enabled, tools have minimal descriptions
|
|
9
|
+
with hints like "See fac://tools/create_document for usage guide."
|
|
10
|
+
|
|
11
|
+
This module provides the get_skill tool that allows LLMs to fetch
|
|
12
|
+
the full documentation on-demand, reducing context token usage by ~90%.
|
|
13
|
+
|
|
14
|
+
Usage:
|
|
15
|
+
from assistant_runtime_sdk import AssistantRuntimeClient
|
|
16
|
+
from assistant_runtime_sdk.skills import SkillProvider
|
|
17
|
+
|
|
18
|
+
client = AssistantRuntimeClient(tenant_id, tenant_secret)
|
|
19
|
+
skills = SkillProvider(client, user_id="user@example.com")
|
|
20
|
+
|
|
21
|
+
# Add to Strands agent tools
|
|
22
|
+
from strands import Agent
|
|
23
|
+
agent = Agent(
|
|
24
|
+
tools=[*mcp_tools, skills.get_skill_tool()]
|
|
25
|
+
)
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
from typing import Dict, Optional, TYPE_CHECKING, Callable
|
|
29
|
+
|
|
30
|
+
if TYPE_CHECKING:
|
|
31
|
+
from .client import AssistantRuntimeClient
|
|
32
|
+
from .async_client import AsyncAssistantRuntimeClient
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class SkillProvider:
|
|
36
|
+
"""
|
|
37
|
+
Provides the get_skill tool for Strands agents.
|
|
38
|
+
|
|
39
|
+
The get_skill tool allows LLMs to fetch detailed documentation for any tool
|
|
40
|
+
before using it. This enables a pattern where tools have minimal descriptions
|
|
41
|
+
(saving context tokens) and the LLM can request full documentation on-demand.
|
|
42
|
+
|
|
43
|
+
Attributes:
|
|
44
|
+
_client: Assistant Runtime client instance for API calls
|
|
45
|
+
_user_id: User identifier for API authentication
|
|
46
|
+
_cache: Optional in-memory cache for skill documentation
|
|
47
|
+
|
|
48
|
+
Example:
|
|
49
|
+
>>> from assistant_runtime_sdk import AssistantRuntimeClient
|
|
50
|
+
>>> from assistant_runtime_sdk.skills import SkillProvider
|
|
51
|
+
>>>
|
|
52
|
+
>>> client = AssistantRuntimeClient(tenant_id, tenant_secret)
|
|
53
|
+
>>> skills = SkillProvider(client, user_id="user@example.com")
|
|
54
|
+
>>>
|
|
55
|
+
>>> # Get the tool function to add to agent
|
|
56
|
+
>>> skill_tool = skills.get_skill_tool()
|
|
57
|
+
>>>
|
|
58
|
+
>>> # Or use directly
|
|
59
|
+
>>> docs = skills.fetch_skill("create_document")
|
|
60
|
+
>>> print(docs)
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
def __init__(
|
|
64
|
+
self,
|
|
65
|
+
client: "AssistantRuntimeClient",
|
|
66
|
+
user_id: str,
|
|
67
|
+
cache_enabled: bool = True,
|
|
68
|
+
):
|
|
69
|
+
"""
|
|
70
|
+
Initialize the SkillProvider.
|
|
71
|
+
|
|
72
|
+
Args:
|
|
73
|
+
client: Assistant Runtime client instance (sync or async)
|
|
74
|
+
user_id: User identifier for API authentication
|
|
75
|
+
cache_enabled: Whether to cache skill documentation (default: True)
|
|
76
|
+
"""
|
|
77
|
+
self._client = client
|
|
78
|
+
self._user_id = user_id
|
|
79
|
+
self._cache: Optional[Dict[str, str]] = {} if cache_enabled else None
|
|
80
|
+
|
|
81
|
+
def get_skill_tool(self) -> Callable[[str], str]:
|
|
82
|
+
"""
|
|
83
|
+
Returns the get_skill tool function decorated for Strands.
|
|
84
|
+
|
|
85
|
+
The returned function can be passed directly to the Strands Agent's
|
|
86
|
+
tools parameter. It will be available to the LLM as a callable tool.
|
|
87
|
+
|
|
88
|
+
Returns:
|
|
89
|
+
A Strands-compatible tool function
|
|
90
|
+
|
|
91
|
+
Example:
|
|
92
|
+
>>> skills = SkillProvider(client, user_id)
|
|
93
|
+
>>> agent = Agent(tools=[skills.get_skill_tool()])
|
|
94
|
+
"""
|
|
95
|
+
try:
|
|
96
|
+
from strands import tool
|
|
97
|
+
except ImportError:
|
|
98
|
+
raise ImportError(
|
|
99
|
+
"strands-agents is required for get_skill_tool(). "
|
|
100
|
+
"Install it with: pip install strands-agents"
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
@tool
|
|
104
|
+
def get_skill(tool_name: str) -> str:
|
|
105
|
+
"""Get detailed documentation for a tool before using it.
|
|
106
|
+
|
|
107
|
+
Call this when you see a tool with a minimal description that says
|
|
108
|
+
"See fac://tools/... for usage guide" to get the full documentation.
|
|
109
|
+
|
|
110
|
+
Args:
|
|
111
|
+
tool_name: The name of the tool (e.g., "create_document", "list_documents")
|
|
112
|
+
|
|
113
|
+
Returns:
|
|
114
|
+
Detailed markdown documentation including:
|
|
115
|
+
- Full description of what the tool does
|
|
116
|
+
- All parameters with types and descriptions
|
|
117
|
+
- Usage examples
|
|
118
|
+
- Important notes and warnings
|
|
119
|
+
"""
|
|
120
|
+
return self.fetch_skill(tool_name)
|
|
121
|
+
|
|
122
|
+
return get_skill
|
|
123
|
+
|
|
124
|
+
def fetch_skill(self, tool_name: str) -> str:
|
|
125
|
+
"""
|
|
126
|
+
Fetch skill documentation for a tool (sync).
|
|
127
|
+
|
|
128
|
+
This is the core method that retrieves documentation from the
|
|
129
|
+
resources API. Results are cached if caching is enabled.
|
|
130
|
+
|
|
131
|
+
Args:
|
|
132
|
+
tool_name: Name of the tool to get documentation for
|
|
133
|
+
|
|
134
|
+
Returns:
|
|
135
|
+
Markdown documentation string, or error message if not found
|
|
136
|
+
"""
|
|
137
|
+
# Check cache first
|
|
138
|
+
if self._cache is not None and tool_name in self._cache:
|
|
139
|
+
return self._cache[tool_name]
|
|
140
|
+
|
|
141
|
+
# Build URI - the MCP server uses fac://tools/{tool_name} format
|
|
142
|
+
uri = f"fac://tools/{tool_name}"
|
|
143
|
+
|
|
144
|
+
try:
|
|
145
|
+
result = self._client.read_resource(self._user_id, uri)
|
|
146
|
+
|
|
147
|
+
if result and result.get("content"):
|
|
148
|
+
content = result["content"]
|
|
149
|
+
|
|
150
|
+
# Cache the result
|
|
151
|
+
if self._cache is not None:
|
|
152
|
+
self._cache[tool_name] = content
|
|
153
|
+
|
|
154
|
+
return content
|
|
155
|
+
|
|
156
|
+
# Check for error response
|
|
157
|
+
if result and result.get("error"):
|
|
158
|
+
return f"Error fetching documentation for '{tool_name}': {result.get('error')}"
|
|
159
|
+
|
|
160
|
+
except Exception as e:
|
|
161
|
+
return f"Could not retrieve documentation for '{tool_name}': {str(e)}"
|
|
162
|
+
|
|
163
|
+
return (
|
|
164
|
+
f"No documentation available for tool '{tool_name}'. "
|
|
165
|
+
"The tool may not have detailed docs, or the resources feature may be disabled on the MCP server."
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
def clear_cache(self) -> None:
|
|
169
|
+
"""Clear the skill documentation cache."""
|
|
170
|
+
if self._cache is not None:
|
|
171
|
+
self._cache.clear()
|
|
172
|
+
|
|
173
|
+
def preload_skills(self, tool_names: list[str]) -> None:
|
|
174
|
+
"""
|
|
175
|
+
Preload documentation for a list of tools.
|
|
176
|
+
|
|
177
|
+
Useful for warming up the cache with commonly used tools
|
|
178
|
+
before starting a conversation.
|
|
179
|
+
|
|
180
|
+
Args:
|
|
181
|
+
tool_names: List of tool names to preload
|
|
182
|
+
"""
|
|
183
|
+
for name in tool_names:
|
|
184
|
+
self.fetch_skill(name)
|
|
185
|
+
|
|
186
|
+
def get_cached_skills(self) -> list[str]:
|
|
187
|
+
"""
|
|
188
|
+
Get list of tool names currently in cache.
|
|
189
|
+
|
|
190
|
+
Returns:
|
|
191
|
+
List of cached tool names, or empty list if caching disabled
|
|
192
|
+
"""
|
|
193
|
+
if self._cache is not None:
|
|
194
|
+
return list(self._cache.keys())
|
|
195
|
+
return []
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
class AsyncSkillProvider:
|
|
199
|
+
"""
|
|
200
|
+
Async version of SkillProvider for use with AsyncAssistantRuntimeClient.
|
|
201
|
+
|
|
202
|
+
Provides the same functionality as SkillProvider but with async methods
|
|
203
|
+
for use in async contexts.
|
|
204
|
+
|
|
205
|
+
Example:
|
|
206
|
+
>>> async with AsyncAssistantRuntimeClient(tenant_id, tenant_secret) as client:
|
|
207
|
+
... skills = AsyncSkillProvider(client, user_id="user@example.com")
|
|
208
|
+
... skill_tool = skills.get_skill_tool()
|
|
209
|
+
... # Add to async agent
|
|
210
|
+
"""
|
|
211
|
+
|
|
212
|
+
def __init__(
|
|
213
|
+
self,
|
|
214
|
+
client: "AsyncAssistantRuntimeClient",
|
|
215
|
+
user_id: str,
|
|
216
|
+
cache_enabled: bool = True,
|
|
217
|
+
):
|
|
218
|
+
"""
|
|
219
|
+
Initialize the AsyncSkillProvider.
|
|
220
|
+
|
|
221
|
+
Args:
|
|
222
|
+
client: AsyncAssistantRuntimeClient instance
|
|
223
|
+
user_id: User identifier for API authentication
|
|
224
|
+
cache_enabled: Whether to cache skill documentation (default: True)
|
|
225
|
+
"""
|
|
226
|
+
self._client = client
|
|
227
|
+
self._user_id = user_id
|
|
228
|
+
self._cache: Optional[Dict[str, str]] = {} if cache_enabled else None
|
|
229
|
+
|
|
230
|
+
def get_skill_tool(self) -> Callable[[str], str]:
|
|
231
|
+
"""
|
|
232
|
+
Returns the async get_skill tool function decorated for Strands.
|
|
233
|
+
|
|
234
|
+
Returns:
|
|
235
|
+
A Strands-compatible async tool function
|
|
236
|
+
"""
|
|
237
|
+
try:
|
|
238
|
+
from strands import tool
|
|
239
|
+
except ImportError:
|
|
240
|
+
raise ImportError(
|
|
241
|
+
"strands-agents is required for get_skill_tool(). "
|
|
242
|
+
"Install it with: pip install strands-agents"
|
|
243
|
+
)
|
|
244
|
+
|
|
245
|
+
@tool
|
|
246
|
+
async def get_skill(tool_name: str) -> str:
|
|
247
|
+
"""Get detailed documentation for a tool before using it.
|
|
248
|
+
|
|
249
|
+
Call this when you see a tool with a minimal description that says
|
|
250
|
+
"See fac://tools/... for usage guide" to get the full documentation.
|
|
251
|
+
|
|
252
|
+
Args:
|
|
253
|
+
tool_name: The name of the tool (e.g., "create_document", "list_documents")
|
|
254
|
+
|
|
255
|
+
Returns:
|
|
256
|
+
Detailed markdown documentation including:
|
|
257
|
+
- Full description of what the tool does
|
|
258
|
+
- All parameters with types and descriptions
|
|
259
|
+
- Usage examples
|
|
260
|
+
- Important notes and warnings
|
|
261
|
+
"""
|
|
262
|
+
return await self.fetch_skill(tool_name)
|
|
263
|
+
|
|
264
|
+
return get_skill
|
|
265
|
+
|
|
266
|
+
async def fetch_skill(self, tool_name: str) -> str:
|
|
267
|
+
"""
|
|
268
|
+
Fetch skill documentation for a tool (async).
|
|
269
|
+
|
|
270
|
+
Args:
|
|
271
|
+
tool_name: Name of the tool to get documentation for
|
|
272
|
+
|
|
273
|
+
Returns:
|
|
274
|
+
Markdown documentation string, or error message if not found
|
|
275
|
+
"""
|
|
276
|
+
# Check cache first
|
|
277
|
+
if self._cache is not None and tool_name in self._cache:
|
|
278
|
+
return self._cache[tool_name]
|
|
279
|
+
|
|
280
|
+
uri = f"fac://tools/{tool_name}"
|
|
281
|
+
|
|
282
|
+
try:
|
|
283
|
+
result = await self._client.read_resource(self._user_id, uri)
|
|
284
|
+
|
|
285
|
+
if result and result.get("content"):
|
|
286
|
+
content = result["content"]
|
|
287
|
+
|
|
288
|
+
if self._cache is not None:
|
|
289
|
+
self._cache[tool_name] = content
|
|
290
|
+
|
|
291
|
+
return content
|
|
292
|
+
|
|
293
|
+
if result and result.get("error"):
|
|
294
|
+
return f"Error fetching documentation for '{tool_name}': {result.get('error')}"
|
|
295
|
+
|
|
296
|
+
except Exception as e:
|
|
297
|
+
return f"Could not retrieve documentation for '{tool_name}': {str(e)}"
|
|
298
|
+
|
|
299
|
+
return (
|
|
300
|
+
f"No documentation available for tool '{tool_name}'. "
|
|
301
|
+
"The tool may not have detailed docs, or the resources feature may be disabled."
|
|
302
|
+
)
|
|
303
|
+
|
|
304
|
+
def clear_cache(self) -> None:
|
|
305
|
+
"""Clear the skill documentation cache."""
|
|
306
|
+
if self._cache is not None:
|
|
307
|
+
self._cache.clear()
|
|
308
|
+
|
|
309
|
+
async def preload_skills(self, tool_names: list[str]) -> None:
|
|
310
|
+
"""
|
|
311
|
+
Preload documentation for a list of tools (async).
|
|
312
|
+
|
|
313
|
+
Args:
|
|
314
|
+
tool_names: List of tool names to preload
|
|
315
|
+
"""
|
|
316
|
+
for name in tool_names:
|
|
317
|
+
await self.fetch_skill(name)
|
|
318
|
+
|
|
319
|
+
def get_cached_skills(self) -> list[str]:
|
|
320
|
+
"""
|
|
321
|
+
Get list of tool names currently in cache.
|
|
322
|
+
|
|
323
|
+
Returns:
|
|
324
|
+
List of cached tool names, or empty list if caching disabled
|
|
325
|
+
"""
|
|
326
|
+
if self._cache is not None:
|
|
327
|
+
return list(self._cache.keys())
|
|
328
|
+
return []
|
|
@@ -0,0 +1,223 @@
|
|
|
1
|
+
# Assistant Runtime SDK - Streaming Utilities
|
|
2
|
+
# Copyright (C) 2025 Paul Clinton
|
|
3
|
+
# AGPL-3.0 License
|
|
4
|
+
|
|
5
|
+
"""
|
|
6
|
+
Server-Sent Events (SSE) parsing utilities for Assistant Runtime streaming responses.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import json
|
|
10
|
+
from enum import Enum
|
|
11
|
+
from typing import Optional, Dict, Any, Iterator, Tuple
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class SSEEventType(str, Enum):
|
|
15
|
+
"""SSE event types used by Assistant Runtime streaming API."""
|
|
16
|
+
|
|
17
|
+
# Core streaming events
|
|
18
|
+
STREAM_START = "stream_start"
|
|
19
|
+
STREAM_CHUNK = "stream_chunk"
|
|
20
|
+
STREAM_COMPLETE = "stream_complete"
|
|
21
|
+
STREAM_ERROR = "stream_error"
|
|
22
|
+
|
|
23
|
+
# AI response events
|
|
24
|
+
THINKING = "thinking"
|
|
25
|
+
|
|
26
|
+
# Tool execution events
|
|
27
|
+
TOOL_CALL_START = "tool_call_start"
|
|
28
|
+
TOOL_CALL_RESULT = "tool_call_result"
|
|
29
|
+
|
|
30
|
+
# Human-in-the-loop events
|
|
31
|
+
APPROVAL_REQUIRED = "approval_required"
|
|
32
|
+
TOOL_CANCELLED = "tool_cancelled"
|
|
33
|
+
|
|
34
|
+
# Auto-model events
|
|
35
|
+
MODEL_FALLBACK = "model_fallback"
|
|
36
|
+
|
|
37
|
+
# Rate limiting events
|
|
38
|
+
RATE_LIMITED = "rate_limited"
|
|
39
|
+
|
|
40
|
+
# Unknown event (fallback)
|
|
41
|
+
UNKNOWN = "unknown"
|
|
42
|
+
|
|
43
|
+
@classmethod
|
|
44
|
+
def from_string(cls, value: str) -> "SSEEventType":
|
|
45
|
+
"""Convert string to SSEEventType, returning UNKNOWN if not found."""
|
|
46
|
+
try:
|
|
47
|
+
return cls(value)
|
|
48
|
+
except ValueError:
|
|
49
|
+
return cls.UNKNOWN
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def parse_sse_line(line: str) -> Optional[Dict[str, Any]]:
|
|
53
|
+
"""
|
|
54
|
+
Parse a single SSE line.
|
|
55
|
+
|
|
56
|
+
Args:
|
|
57
|
+
line: Raw SSE line (e.g., "event: stream_chunk" or "data: {...}")
|
|
58
|
+
|
|
59
|
+
Returns:
|
|
60
|
+
Parsed event dict with 'type' and 'value', or None for empty/comment lines
|
|
61
|
+
|
|
62
|
+
Example:
|
|
63
|
+
>>> parse_sse_line("event: stream_chunk")
|
|
64
|
+
{'type': 'event_name', 'value': 'stream_chunk'}
|
|
65
|
+
>>> parse_sse_line("data: {\"content\": \"Hello\"}")
|
|
66
|
+
{'type': 'data', 'value': {'content': 'Hello'}}
|
|
67
|
+
>>> parse_sse_line(": heartbeat")
|
|
68
|
+
{'type': 'heartbeat', 'value': 'heartbeat'}
|
|
69
|
+
"""
|
|
70
|
+
line = line.strip()
|
|
71
|
+
|
|
72
|
+
# Skip empty lines
|
|
73
|
+
if not line:
|
|
74
|
+
return None
|
|
75
|
+
|
|
76
|
+
# Surface heartbeat comments so callers can use them as keepalive signals
|
|
77
|
+
if line.startswith(":"):
|
|
78
|
+
return {"type": "heartbeat", "value": line[1:].strip()}
|
|
79
|
+
|
|
80
|
+
if line.startswith("event:"):
|
|
81
|
+
return {"type": "event_name", "value": line[6:].strip()}
|
|
82
|
+
elif line.startswith("data:"):
|
|
83
|
+
data_str = line[5:].strip()
|
|
84
|
+
if data_str:
|
|
85
|
+
try:
|
|
86
|
+
return {"type": "data", "value": json.loads(data_str)}
|
|
87
|
+
except json.JSONDecodeError:
|
|
88
|
+
# Return raw string if not valid JSON
|
|
89
|
+
return {"type": "data", "value": data_str}
|
|
90
|
+
elif line.startswith("id:"):
|
|
91
|
+
return {"type": "id", "value": line[3:].strip()}
|
|
92
|
+
elif line.startswith("retry:"):
|
|
93
|
+
try:
|
|
94
|
+
return {"type": "retry", "value": int(line[6:].strip())}
|
|
95
|
+
except ValueError:
|
|
96
|
+
return None
|
|
97
|
+
|
|
98
|
+
return None
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def parse_sse_stream(lines: Iterator[str]) -> Iterator[Dict[str, Any]]:
|
|
102
|
+
"""
|
|
103
|
+
Parse an SSE stream from an iterator of lines.
|
|
104
|
+
|
|
105
|
+
Yields complete events as dicts with 'event' and 'data' keys.
|
|
106
|
+
|
|
107
|
+
Args:
|
|
108
|
+
lines: Iterator of SSE lines (from response.iter_lines() or similar)
|
|
109
|
+
|
|
110
|
+
Yields:
|
|
111
|
+
Event dicts: {'event': 'stream_chunk', 'data': {...}}
|
|
112
|
+
|
|
113
|
+
Example:
|
|
114
|
+
>>> for event in parse_sse_stream(response.iter_lines()):
|
|
115
|
+
... print(event['event'], event['data'])
|
|
116
|
+
"""
|
|
117
|
+
current_event: Optional[str] = None
|
|
118
|
+
|
|
119
|
+
for line in lines:
|
|
120
|
+
if line is None:
|
|
121
|
+
continue
|
|
122
|
+
|
|
123
|
+
# Handle both str and bytes
|
|
124
|
+
if isinstance(line, bytes):
|
|
125
|
+
line = line.decode("utf-8")
|
|
126
|
+
|
|
127
|
+
parsed = parse_sse_line(line)
|
|
128
|
+
if not parsed:
|
|
129
|
+
continue
|
|
130
|
+
|
|
131
|
+
if parsed["type"] == "heartbeat":
|
|
132
|
+
yield {"event": "heartbeat", "data": {}}
|
|
133
|
+
elif parsed["type"] == "event_name":
|
|
134
|
+
current_event = parsed["value"]
|
|
135
|
+
elif parsed["type"] == "data":
|
|
136
|
+
yield {"event": current_event or "message", "data": parsed["value"]}
|
|
137
|
+
current_event = None
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def format_sse_event(event_type: str, data: Any, event_id: Optional[str] = None) -> str:
|
|
141
|
+
"""
|
|
142
|
+
Format data as an SSE event string.
|
|
143
|
+
|
|
144
|
+
Useful for servers that need to send SSE responses.
|
|
145
|
+
|
|
146
|
+
Args:
|
|
147
|
+
event_type: Event type name
|
|
148
|
+
data: Event data (will be JSON serialized)
|
|
149
|
+
event_id: Optional event ID
|
|
150
|
+
|
|
151
|
+
Returns:
|
|
152
|
+
Formatted SSE event string
|
|
153
|
+
|
|
154
|
+
Example:
|
|
155
|
+
>>> print(format_sse_event("stream_chunk", {"content": "Hello"}))
|
|
156
|
+
event: stream_chunk
|
|
157
|
+
data: {"content": "Hello"}
|
|
158
|
+
|
|
159
|
+
"""
|
|
160
|
+
lines = []
|
|
161
|
+
|
|
162
|
+
if event_id:
|
|
163
|
+
lines.append(f"id: {event_id}")
|
|
164
|
+
|
|
165
|
+
lines.append(f"event: {event_type}")
|
|
166
|
+
|
|
167
|
+
if isinstance(data, str):
|
|
168
|
+
lines.append(f"data: {data}")
|
|
169
|
+
else:
|
|
170
|
+
lines.append(f"data: {json.dumps(data)}")
|
|
171
|
+
|
|
172
|
+
lines.append("") # Empty line to end event
|
|
173
|
+
return "\n".join(lines) + "\n"
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def is_terminal_event(event_type: str) -> bool:
|
|
177
|
+
"""
|
|
178
|
+
Check if an event type indicates the stream has ended.
|
|
179
|
+
|
|
180
|
+
Args:
|
|
181
|
+
event_type: The event type string
|
|
182
|
+
|
|
183
|
+
Returns:
|
|
184
|
+
True if this is a terminal event (stream_complete, stream_error, rate_limited)
|
|
185
|
+
|
|
186
|
+
Example:
|
|
187
|
+
>>> is_terminal_event("stream_chunk")
|
|
188
|
+
False
|
|
189
|
+
>>> is_terminal_event("stream_complete")
|
|
190
|
+
True
|
|
191
|
+
"""
|
|
192
|
+
terminal_events = {
|
|
193
|
+
SSEEventType.STREAM_COMPLETE.value,
|
|
194
|
+
SSEEventType.STREAM_ERROR.value,
|
|
195
|
+
SSEEventType.RATE_LIMITED.value,
|
|
196
|
+
}
|
|
197
|
+
return event_type in terminal_events
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def extract_stream_metrics(complete_event: Dict[str, Any]) -> Dict[str, Any]:
|
|
201
|
+
"""
|
|
202
|
+
Extract metrics from a stream_complete event.
|
|
203
|
+
|
|
204
|
+
Args:
|
|
205
|
+
complete_event: The data from a stream_complete event
|
|
206
|
+
|
|
207
|
+
Returns:
|
|
208
|
+
Dict with normalized metrics (tokens_used, model_id, duration_ms, etc.)
|
|
209
|
+
|
|
210
|
+
Example:
|
|
211
|
+
>>> metrics = extract_stream_metrics(event['data'])
|
|
212
|
+
>>> print(f"Used {metrics.get('tokens_used', 0)} tokens")
|
|
213
|
+
"""
|
|
214
|
+
data = complete_event.get("data", complete_event) if isinstance(complete_event, dict) else {}
|
|
215
|
+
|
|
216
|
+
return {
|
|
217
|
+
"tokens_used": data.get("tokens_used", 0),
|
|
218
|
+
"tokens_actual": data.get("tokens_actual", 0),
|
|
219
|
+
"model_id": data.get("model_id"),
|
|
220
|
+
"message_id": data.get("message_id"),
|
|
221
|
+
"session_id": data.get("session_id"),
|
|
222
|
+
"full_response": data.get("full_response", ""),
|
|
223
|
+
}
|