metergraph 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.
- metergraph/__init__.py +219 -0
- metergraph/_capture.py +1146 -0
- metergraph/_config.py +144 -0
- metergraph/_context.py +136 -0
- metergraph/_template.py +58 -0
- metergraph/_track.py +61 -0
- metergraph/_transport.py +171 -0
- metergraph/_version.py +11 -0
- metergraph-0.1.0.dist-info/METADATA +88 -0
- metergraph-0.1.0.dist-info/RECORD +12 -0
- metergraph-0.1.0.dist-info/WHEEL +5 -0
- metergraph-0.1.0.dist-info/top_level.txt +1 -0
metergraph/_capture.py
ADDED
|
@@ -0,0 +1,1146 @@
|
|
|
1
|
+
"""Provider-client wrapping and normalized record construction."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import functools
|
|
6
|
+
import inspect
|
|
7
|
+
import json
|
|
8
|
+
import logging
|
|
9
|
+
import os
|
|
10
|
+
import platform
|
|
11
|
+
import sys
|
|
12
|
+
import time
|
|
13
|
+
from dataclasses import dataclass
|
|
14
|
+
from datetime import datetime, timezone
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
from typing import Any, Callable, Mapping
|
|
17
|
+
|
|
18
|
+
from ._context import CaptureContext, snapshot
|
|
19
|
+
from ._template import scrub, template_hash
|
|
20
|
+
from ._version import SDK_VERSION
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
log = logging.getLogger("metergraph")
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def _get(value: Any, name: str, default: Any = None) -> Any:
|
|
27
|
+
if isinstance(value, Mapping):
|
|
28
|
+
return value.get(name, default)
|
|
29
|
+
return getattr(value, name, default)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _first(value: Any) -> Any:
|
|
33
|
+
try:
|
|
34
|
+
return value[0]
|
|
35
|
+
except (IndexError, KeyError, TypeError):
|
|
36
|
+
return None
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _int(value: Any) -> int | None:
|
|
40
|
+
try:
|
|
41
|
+
return int(value) if value is not None else None
|
|
42
|
+
except (TypeError, ValueError):
|
|
43
|
+
return None
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _usage(response: Any) -> dict[str, int | None]:
|
|
47
|
+
usage = (
|
|
48
|
+
_get(response, "usage")
|
|
49
|
+
or _get(response, "usage_metadata")
|
|
50
|
+
or _get(response, "usageMetadata")
|
|
51
|
+
)
|
|
52
|
+
prompt_details = _get(usage, "prompt_tokens_details") or _get(
|
|
53
|
+
usage, "input_tokens_details"
|
|
54
|
+
) or _get(
|
|
55
|
+
usage, "promptTokensDetails"
|
|
56
|
+
) or _get(
|
|
57
|
+
usage, "inputTokensDetails"
|
|
58
|
+
)
|
|
59
|
+
completion_details = _get(usage, "completion_tokens_details") or _get(
|
|
60
|
+
usage, "output_tokens_details"
|
|
61
|
+
) or _get(
|
|
62
|
+
usage, "completionTokensDetails"
|
|
63
|
+
) or _get(
|
|
64
|
+
usage, "outputTokensDetails"
|
|
65
|
+
)
|
|
66
|
+
cache_creation = _get(usage, "cache_creation") or _get(
|
|
67
|
+
usage, "cacheCreation"
|
|
68
|
+
)
|
|
69
|
+
return {
|
|
70
|
+
"input_tokens": _int(
|
|
71
|
+
_get(
|
|
72
|
+
usage,
|
|
73
|
+
"prompt_tokens",
|
|
74
|
+
_get(
|
|
75
|
+
usage,
|
|
76
|
+
"input_tokens",
|
|
77
|
+
_get(
|
|
78
|
+
usage,
|
|
79
|
+
"prompt_token_count",
|
|
80
|
+
_get(usage, "promptTokenCount"),
|
|
81
|
+
),
|
|
82
|
+
),
|
|
83
|
+
)
|
|
84
|
+
),
|
|
85
|
+
"output_tokens": _int(
|
|
86
|
+
_get(
|
|
87
|
+
usage,
|
|
88
|
+
"completion_tokens",
|
|
89
|
+
_get(
|
|
90
|
+
usage,
|
|
91
|
+
"output_tokens",
|
|
92
|
+
_get(
|
|
93
|
+
usage,
|
|
94
|
+
"candidates_token_count",
|
|
95
|
+
_get(usage, "candidatesTokenCount"),
|
|
96
|
+
),
|
|
97
|
+
),
|
|
98
|
+
)
|
|
99
|
+
),
|
|
100
|
+
"cache_read_tokens": _int(
|
|
101
|
+
_get(
|
|
102
|
+
usage,
|
|
103
|
+
"cache_read_input_tokens",
|
|
104
|
+
_get(
|
|
105
|
+
prompt_details,
|
|
106
|
+
"cached_tokens",
|
|
107
|
+
_get(
|
|
108
|
+
usage,
|
|
109
|
+
"cached_content_token_count",
|
|
110
|
+
_get(usage, "cachedContentTokenCount"),
|
|
111
|
+
),
|
|
112
|
+
),
|
|
113
|
+
)
|
|
114
|
+
),
|
|
115
|
+
"cache_write_tokens": _int(
|
|
116
|
+
_get(
|
|
117
|
+
usage,
|
|
118
|
+
"cache_creation_input_tokens",
|
|
119
|
+
_get(
|
|
120
|
+
prompt_details,
|
|
121
|
+
"cache_write_tokens",
|
|
122
|
+
_get(
|
|
123
|
+
usage,
|
|
124
|
+
"cacheCreationInputTokens",
|
|
125
|
+
_get(prompt_details, "cacheWriteTokens"),
|
|
126
|
+
),
|
|
127
|
+
),
|
|
128
|
+
)
|
|
129
|
+
),
|
|
130
|
+
"cache_write_5m_tokens": _int(
|
|
131
|
+
_get(
|
|
132
|
+
cache_creation,
|
|
133
|
+
"ephemeral_5m_input_tokens",
|
|
134
|
+
_get(cache_creation, "ephemeral5mInputTokens"),
|
|
135
|
+
)
|
|
136
|
+
),
|
|
137
|
+
"cache_write_1h_tokens": _int(
|
|
138
|
+
_get(
|
|
139
|
+
cache_creation,
|
|
140
|
+
"ephemeral_1h_input_tokens",
|
|
141
|
+
_get(cache_creation, "ephemeral1hInputTokens"),
|
|
142
|
+
)
|
|
143
|
+
),
|
|
144
|
+
"reasoning_tokens": _int(
|
|
145
|
+
_get(
|
|
146
|
+
completion_details,
|
|
147
|
+
"reasoning_tokens",
|
|
148
|
+
_get(
|
|
149
|
+
usage,
|
|
150
|
+
"thoughts_token_count",
|
|
151
|
+
_get(usage, "thoughtsTokenCount"),
|
|
152
|
+
),
|
|
153
|
+
)
|
|
154
|
+
),
|
|
155
|
+
}
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def _response_text(response: Any) -> str | None:
|
|
159
|
+
direct = _get(response, "output_text") or _get(response, "text")
|
|
160
|
+
if isinstance(direct, str):
|
|
161
|
+
return direct
|
|
162
|
+
choice = _first(_get(response, "choices"))
|
|
163
|
+
message = _get(choice, "message")
|
|
164
|
+
content = _get(message, "content")
|
|
165
|
+
if isinstance(content, str):
|
|
166
|
+
return content
|
|
167
|
+
blocks = _get(response, "content")
|
|
168
|
+
if isinstance(blocks, list):
|
|
169
|
+
texts = [_get(block, "text") for block in blocks]
|
|
170
|
+
joined = "".join(text for text in texts if isinstance(text, str))
|
|
171
|
+
if joined:
|
|
172
|
+
return joined
|
|
173
|
+
outputs = _get(response, "output")
|
|
174
|
+
if isinstance(outputs, list):
|
|
175
|
+
texts = []
|
|
176
|
+
for output in outputs:
|
|
177
|
+
content = _get(output, "content")
|
|
178
|
+
if not isinstance(content, list):
|
|
179
|
+
continue
|
|
180
|
+
for block in content:
|
|
181
|
+
text = _get(block, "text") or _get(block, "output_text")
|
|
182
|
+
if isinstance(text, str):
|
|
183
|
+
texts.append(text)
|
|
184
|
+
return "".join(texts) or None
|
|
185
|
+
return None
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def _chunk_text(chunk: Any) -> str | None:
|
|
189
|
+
choice = _first(_get(chunk, "choices"))
|
|
190
|
+
delta = _get(choice, "delta")
|
|
191
|
+
content = _get(delta, "content")
|
|
192
|
+
if isinstance(content, str):
|
|
193
|
+
return content
|
|
194
|
+
delta = _get(chunk, "delta")
|
|
195
|
+
text = _get(delta, "text") or _get(chunk, "text")
|
|
196
|
+
if isinstance(text, str):
|
|
197
|
+
return text
|
|
198
|
+
if _get(chunk, "type") == "content_block_delta":
|
|
199
|
+
text = _get(_get(chunk, "delta"), "text")
|
|
200
|
+
return text if isinstance(text, str) else None
|
|
201
|
+
return None
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
def _usage_only_chunk(chunk: Any, call: "CallState") -> bool:
|
|
205
|
+
return (
|
|
206
|
+
call.provider == "openai"
|
|
207
|
+
and call.endpoint == "chat.completions"
|
|
208
|
+
and _get(chunk, "choices") == []
|
|
209
|
+
and _get(chunk, "usage") is not None
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
def _stop_reason(response: Any) -> str | None:
|
|
214
|
+
direct = _get(response, "stop_reason") or _get(response, "status")
|
|
215
|
+
if direct:
|
|
216
|
+
return str(direct)
|
|
217
|
+
choice = _first(_get(response, "choices"))
|
|
218
|
+
reason = _get(choice, "finish_reason")
|
|
219
|
+
return str(reason) if reason is not None else None
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _request_id(response: Any) -> str | None:
|
|
223
|
+
value = (
|
|
224
|
+
_get(response, "_request_id")
|
|
225
|
+
or _get(response, "response_id")
|
|
226
|
+
or _get(response, "id")
|
|
227
|
+
)
|
|
228
|
+
return str(value) if value is not None else None
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def _tool_names(request: Mapping[str, Any]) -> list[dict[str, str]] | None:
|
|
232
|
+
tools = request.get("tools")
|
|
233
|
+
if not isinstance(tools, list):
|
|
234
|
+
return None
|
|
235
|
+
names: list[dict[str, str]] = []
|
|
236
|
+
for tool in tools:
|
|
237
|
+
fn = _get(tool, "function")
|
|
238
|
+
name = _get(fn, "name") or _get(tool, "name")
|
|
239
|
+
if name:
|
|
240
|
+
names.append({"name": str(name)})
|
|
241
|
+
return names or None
|
|
242
|
+
|
|
243
|
+
|
|
244
|
+
def _tool_argument(value: Any) -> Any:
|
|
245
|
+
if isinstance(value, str):
|
|
246
|
+
try:
|
|
247
|
+
return json.loads(value)
|
|
248
|
+
except json.JSONDecodeError:
|
|
249
|
+
return value
|
|
250
|
+
return scrub(value)
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def _tool_policies(request: Mapping[str, Any]) -> dict[str, str]:
|
|
254
|
+
policies: dict[str, str] = {}
|
|
255
|
+
tools = request.get("tools")
|
|
256
|
+
if not isinstance(tools, list):
|
|
257
|
+
return policies
|
|
258
|
+
for tool in tools:
|
|
259
|
+
fn = _get(tool, "function")
|
|
260
|
+
name = _get(fn, "name") or _get(tool, "name")
|
|
261
|
+
policy = (
|
|
262
|
+
_get(fn, "x-metergraph-idempotency")
|
|
263
|
+
or _get(tool, "x-metergraph-idempotency")
|
|
264
|
+
or _get(fn, "metergraph_idempotency")
|
|
265
|
+
or _get(tool, "metergraph_idempotency")
|
|
266
|
+
)
|
|
267
|
+
if name:
|
|
268
|
+
policies[str(name)] = (
|
|
269
|
+
str(policy)
|
|
270
|
+
if policy in {"idempotent", "non_idempotent"}
|
|
271
|
+
else "non_idempotent"
|
|
272
|
+
)
|
|
273
|
+
return policies
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def _tool_events(request: Mapping[str, Any], response: Any) -> list[dict] | None:
|
|
277
|
+
"""Normalize completed history and newly requested provider tool calls."""
|
|
278
|
+
policies = _tool_policies(request)
|
|
279
|
+
calls: dict[str, dict] = {}
|
|
280
|
+
order: list[str] = []
|
|
281
|
+
pending_results: dict[str, tuple[Any, bool]] = {}
|
|
282
|
+
|
|
283
|
+
def call(call_id: Any, name: Any, arguments: Any) -> None:
|
|
284
|
+
if not name:
|
|
285
|
+
return
|
|
286
|
+
key = str(call_id or f"{name}:{len(order)}")
|
|
287
|
+
if key not in calls:
|
|
288
|
+
order.append(key)
|
|
289
|
+
calls[key] = {
|
|
290
|
+
"call_id": key,
|
|
291
|
+
"name": str(name),
|
|
292
|
+
"arguments": _tool_argument(arguments),
|
|
293
|
+
"result": None,
|
|
294
|
+
"status": "requested",
|
|
295
|
+
"idempotency": policies.get(str(name), "non_idempotent"),
|
|
296
|
+
}
|
|
297
|
+
if key in pending_results:
|
|
298
|
+
result, is_error = pending_results.pop(key)
|
|
299
|
+
complete(key, result, is_error)
|
|
300
|
+
|
|
301
|
+
def complete(call_id: Any, result: Any, is_error: bool = False) -> None:
|
|
302
|
+
key = str(call_id or "")
|
|
303
|
+
if key not in calls:
|
|
304
|
+
pending_results[key] = (result, is_error)
|
|
305
|
+
return
|
|
306
|
+
calls[key]["result"] = _tool_argument(result)
|
|
307
|
+
calls[key]["status"] = "error" if is_error else "completed"
|
|
308
|
+
|
|
309
|
+
def content_blocks(value: Any) -> None:
|
|
310
|
+
if not isinstance(value, list):
|
|
311
|
+
return
|
|
312
|
+
for block in value:
|
|
313
|
+
kind = _get(block, "type")
|
|
314
|
+
if kind == "tool_use":
|
|
315
|
+
call(_get(block, "id"), _get(block, "name"), _get(block, "input"))
|
|
316
|
+
elif kind == "tool_result":
|
|
317
|
+
complete(
|
|
318
|
+
_get(block, "tool_use_id"),
|
|
319
|
+
_get(block, "content"),
|
|
320
|
+
bool(_get(block, "is_error", False)),
|
|
321
|
+
)
|
|
322
|
+
|
|
323
|
+
history = request.get("messages")
|
|
324
|
+
if not isinstance(history, list):
|
|
325
|
+
history = request.get("input")
|
|
326
|
+
if isinstance(history, list):
|
|
327
|
+
for message in history:
|
|
328
|
+
for tool_call in _get(message, "tool_calls", []) or []:
|
|
329
|
+
fn = _get(tool_call, "function")
|
|
330
|
+
call(
|
|
331
|
+
_get(tool_call, "id") or _get(tool_call, "call_id"),
|
|
332
|
+
_get(fn, "name") or _get(tool_call, "name"),
|
|
333
|
+
_get(fn, "arguments", _get(tool_call, "arguments")),
|
|
334
|
+
)
|
|
335
|
+
role = _get(message, "role")
|
|
336
|
+
if role == "tool":
|
|
337
|
+
complete(
|
|
338
|
+
_get(message, "tool_call_id") or _get(message, "call_id"),
|
|
339
|
+
_get(message, "content"),
|
|
340
|
+
bool(_get(message, "is_error", False)),
|
|
341
|
+
)
|
|
342
|
+
kind = _get(message, "type")
|
|
343
|
+
if kind in {"function_call", "tool_call"}:
|
|
344
|
+
call(
|
|
345
|
+
_get(message, "call_id") or _get(message, "id"),
|
|
346
|
+
_get(message, "name"),
|
|
347
|
+
_get(message, "arguments"),
|
|
348
|
+
)
|
|
349
|
+
elif kind in {"function_call_output", "tool_result"}:
|
|
350
|
+
complete(
|
|
351
|
+
_get(message, "call_id") or _get(message, "tool_use_id"),
|
|
352
|
+
_get(message, "output", _get(message, "content")),
|
|
353
|
+
bool(_get(message, "is_error", False)),
|
|
354
|
+
)
|
|
355
|
+
content_blocks(_get(message, "content"))
|
|
356
|
+
|
|
357
|
+
choice = _first(_get(response, "choices"))
|
|
358
|
+
response_message = _get(choice, "message")
|
|
359
|
+
for tool_call in _get(response_message, "tool_calls", []) or []:
|
|
360
|
+
fn = _get(tool_call, "function")
|
|
361
|
+
call(
|
|
362
|
+
_get(tool_call, "id") or _get(tool_call, "call_id"),
|
|
363
|
+
_get(fn, "name") or _get(tool_call, "name"),
|
|
364
|
+
_get(fn, "arguments", _get(tool_call, "arguments")),
|
|
365
|
+
)
|
|
366
|
+
content_blocks(_get(response, "content"))
|
|
367
|
+
for output in _get(response, "output", []) or []:
|
|
368
|
+
kind = _get(output, "type")
|
|
369
|
+
if kind in {"function_call", "tool_call"}:
|
|
370
|
+
call(
|
|
371
|
+
_get(output, "call_id") or _get(output, "id"),
|
|
372
|
+
_get(output, "name"),
|
|
373
|
+
_get(output, "arguments"),
|
|
374
|
+
)
|
|
375
|
+
|
|
376
|
+
return [calls[key] for key in order] or None
|
|
377
|
+
|
|
378
|
+
|
|
379
|
+
def _capture_frames(
|
|
380
|
+
app_root: str, skip_frames: tuple[str, ...]
|
|
381
|
+
) -> tuple[str | None, str | None, list[dict]]:
|
|
382
|
+
frames: list[dict] = []
|
|
383
|
+
root = os.path.realpath(app_root)
|
|
384
|
+
frame = sys._getframe(2)
|
|
385
|
+
while frame is not None and len(frames) < 5:
|
|
386
|
+
filename = os.path.realpath(frame.f_code.co_filename)
|
|
387
|
+
if filename.startswith(root) and not any(
|
|
388
|
+
part in filename for part in skip_frames
|
|
389
|
+
):
|
|
390
|
+
relative = os.path.relpath(filename, root)
|
|
391
|
+
module = str(Path(relative).with_suffix("")).replace(os.sep, ".")
|
|
392
|
+
qualname = getattr(frame.f_code, "co_qualname", frame.f_code.co_name)
|
|
393
|
+
frames.append({"m": module, "f": qualname, "l": frame.f_lineno})
|
|
394
|
+
frame = frame.f_back
|
|
395
|
+
if not frames:
|
|
396
|
+
return None, None, []
|
|
397
|
+
return f"{frames[0]['m']}:{frames[0]['f']}", frames[0]["m"], frames
|
|
398
|
+
|
|
399
|
+
|
|
400
|
+
@dataclass
|
|
401
|
+
class Options:
|
|
402
|
+
capture_text: bool = False
|
|
403
|
+
redact: Callable[[str, str], str] | None = None
|
|
404
|
+
app_root: str = os.getcwd()
|
|
405
|
+
skip_frames: tuple[str, ...] = ()
|
|
406
|
+
environment: str | None = None
|
|
407
|
+
text_max_bytes: int = 100_000
|
|
408
|
+
|
|
409
|
+
|
|
410
|
+
class Runtime:
|
|
411
|
+
def __init__(self, writer: Any, options: Options) -> None:
|
|
412
|
+
self.writer = writer
|
|
413
|
+
self.options = options
|
|
414
|
+
|
|
415
|
+
def call_state(
|
|
416
|
+
self,
|
|
417
|
+
provider: str,
|
|
418
|
+
endpoint: str,
|
|
419
|
+
request: Mapping[str, Any],
|
|
420
|
+
*,
|
|
421
|
+
context: CaptureContext | None = None,
|
|
422
|
+
) -> "CallState":
|
|
423
|
+
context = context or snapshot()
|
|
424
|
+
func, module, frames = _capture_frames(
|
|
425
|
+
self.options.app_root,
|
|
426
|
+
(
|
|
427
|
+
"site-packages",
|
|
428
|
+
"metergraph/_capture.py",
|
|
429
|
+
"concurrent/futures",
|
|
430
|
+
"threading.py",
|
|
431
|
+
*self.options.skip_frames,
|
|
432
|
+
),
|
|
433
|
+
)
|
|
434
|
+
return CallState(
|
|
435
|
+
runtime=self,
|
|
436
|
+
provider=provider,
|
|
437
|
+
endpoint=endpoint,
|
|
438
|
+
request=dict(request),
|
|
439
|
+
context=context,
|
|
440
|
+
started=time.perf_counter(),
|
|
441
|
+
ts=datetime.now(timezone.utc).isoformat(),
|
|
442
|
+
func=context.func_name or func,
|
|
443
|
+
module=context.func_module or module,
|
|
444
|
+
frames=frames,
|
|
445
|
+
)
|
|
446
|
+
|
|
447
|
+
def _text(
|
|
448
|
+
self, value: str | None, kind: str, *, enabled: bool
|
|
449
|
+
) -> tuple[str | None, bool]:
|
|
450
|
+
if not enabled or value is None:
|
|
451
|
+
return None, False
|
|
452
|
+
if self.options.redact:
|
|
453
|
+
try:
|
|
454
|
+
value = self.options.redact(value, kind)
|
|
455
|
+
except Exception:
|
|
456
|
+
return "<redaction-failed>", False
|
|
457
|
+
raw = value.encode()
|
|
458
|
+
if len(raw) <= self.options.text_max_bytes:
|
|
459
|
+
return value, False
|
|
460
|
+
marker = "\n<metergraph:truncated>"
|
|
461
|
+
clipped = raw[: max(0, self.options.text_max_bytes - len(marker))].decode(
|
|
462
|
+
errors="ignore"
|
|
463
|
+
)
|
|
464
|
+
return clipped + marker, True
|
|
465
|
+
|
|
466
|
+
|
|
467
|
+
@dataclass
|
|
468
|
+
class CallState:
|
|
469
|
+
runtime: Runtime
|
|
470
|
+
provider: str
|
|
471
|
+
endpoint: str
|
|
472
|
+
request: dict[str, Any]
|
|
473
|
+
context: CaptureContext
|
|
474
|
+
started: float
|
|
475
|
+
ts: str
|
|
476
|
+
func: str | None
|
|
477
|
+
module: str | None
|
|
478
|
+
frames: list[dict]
|
|
479
|
+
done: bool = False
|
|
480
|
+
|
|
481
|
+
def finish(
|
|
482
|
+
self,
|
|
483
|
+
response: Any = None,
|
|
484
|
+
*,
|
|
485
|
+
status: str | None = None,
|
|
486
|
+
error: BaseException | None = None,
|
|
487
|
+
stream: bool = False,
|
|
488
|
+
ttft_ms: int | None = None,
|
|
489
|
+
response_text: str | None = None,
|
|
490
|
+
) -> None:
|
|
491
|
+
if self.done:
|
|
492
|
+
return
|
|
493
|
+
self.done = True
|
|
494
|
+
capture_text = (
|
|
495
|
+
self.context.capture_text
|
|
496
|
+
if self.context.capture_text is not None
|
|
497
|
+
else self.runtime.options.capture_text
|
|
498
|
+
)
|
|
499
|
+
request_clean = scrub(self.request)
|
|
500
|
+
request_json, request_truncated = self.runtime._text(
|
|
501
|
+
json.dumps(request_clean, separators=(",", ":"), default=repr),
|
|
502
|
+
"request",
|
|
503
|
+
enabled=capture_text,
|
|
504
|
+
)
|
|
505
|
+
response_json, response_truncated = self.runtime._text(
|
|
506
|
+
response_text if response_text is not None else _response_text(response),
|
|
507
|
+
"response",
|
|
508
|
+
enabled=capture_text,
|
|
509
|
+
)
|
|
510
|
+
tool_calls = _tool_events(request_clean, response)
|
|
511
|
+
tool_truncated = False
|
|
512
|
+
if tool_calls and capture_text:
|
|
513
|
+
encoded_tools, tool_truncated = self.runtime._text(
|
|
514
|
+
json.dumps(tool_calls, separators=(",", ":"), default=repr),
|
|
515
|
+
"tool_calls",
|
|
516
|
+
enabled=True,
|
|
517
|
+
)
|
|
518
|
+
try:
|
|
519
|
+
tool_calls = json.loads(encoded_tools) if encoded_tools else None
|
|
520
|
+
except json.JSONDecodeError:
|
|
521
|
+
tool_calls = None
|
|
522
|
+
elif tool_calls:
|
|
523
|
+
tool_calls = [
|
|
524
|
+
{
|
|
525
|
+
"call_id": item["call_id"],
|
|
526
|
+
"name": item["name"],
|
|
527
|
+
"status": item["status"],
|
|
528
|
+
"idempotency": item["idempotency"],
|
|
529
|
+
}
|
|
530
|
+
for item in tool_calls
|
|
531
|
+
]
|
|
532
|
+
row: dict[str, Any] = {
|
|
533
|
+
"ts": self.ts,
|
|
534
|
+
"route": self.context.route,
|
|
535
|
+
"provider": self.provider,
|
|
536
|
+
"model": self.request.get("model"),
|
|
537
|
+
**_usage(response),
|
|
538
|
+
"latency_ms": round((time.perf_counter() - self.started) * 1000),
|
|
539
|
+
"status": status
|
|
540
|
+
or ("error" if error else _stop_reason(response) or "success"),
|
|
541
|
+
"session_id": self.context.session_id,
|
|
542
|
+
"conversation_id": self.context.session_id,
|
|
543
|
+
"template_hash": template_hash(self.request),
|
|
544
|
+
"unit_name": self.context.unit_name,
|
|
545
|
+
"unit_count": self.context.unit_count,
|
|
546
|
+
"tool_calls": tool_calls,
|
|
547
|
+
"endpoint": self.endpoint,
|
|
548
|
+
"request_id": _request_id(response),
|
|
549
|
+
"batch": self.request.get("batch") is True,
|
|
550
|
+
"batch_custom_id": self.request.get("batch_custom_id"),
|
|
551
|
+
# The worker requires an explicit positive stamp before sending
|
|
552
|
+
# any content to Bedrock. Missing/old rows therefore fail closed.
|
|
553
|
+
"content_opted_in": capture_text,
|
|
554
|
+
"request_json": request_json,
|
|
555
|
+
"response_text": response_json,
|
|
556
|
+
"text_truncated": request_truncated or response_truncated or tool_truncated,
|
|
557
|
+
"stream": stream,
|
|
558
|
+
"ttft_ms": ttft_ms,
|
|
559
|
+
"func": self.func,
|
|
560
|
+
"module": self.module,
|
|
561
|
+
"frames_json": self.frames,
|
|
562
|
+
"tags": dict(self.context.tags),
|
|
563
|
+
"environment": self.runtime.options.environment,
|
|
564
|
+
"error": bool(error),
|
|
565
|
+
"error_type": type(error).__name__ if error else None,
|
|
566
|
+
"sdk": "python",
|
|
567
|
+
"sdk_version": SDK_VERSION,
|
|
568
|
+
"runtime": f"{platform.python_implementation().lower()}-{platform.python_version()}",
|
|
569
|
+
}
|
|
570
|
+
try:
|
|
571
|
+
self.runtime.writer.enqueue(row)
|
|
572
|
+
except Exception:
|
|
573
|
+
pass
|
|
574
|
+
|
|
575
|
+
|
|
576
|
+
class _StreamState:
|
|
577
|
+
def __init__(self, stream: Any, call: CallState) -> None:
|
|
578
|
+
self.stream = stream
|
|
579
|
+
self.call = call
|
|
580
|
+
self.iterator = None
|
|
581
|
+
self.last = None
|
|
582
|
+
self.parts: list[str] = []
|
|
583
|
+
self.ttft_ms: int | None = None
|
|
584
|
+
|
|
585
|
+
def chunk(self, value: Any) -> Any:
|
|
586
|
+
self.last = value
|
|
587
|
+
text = _chunk_text(value)
|
|
588
|
+
if text:
|
|
589
|
+
if self.ttft_ms is None:
|
|
590
|
+
self.ttft_ms = round((time.perf_counter() - self.call.started) * 1000)
|
|
591
|
+
self.parts.append(text)
|
|
592
|
+
return value
|
|
593
|
+
|
|
594
|
+
def finish(
|
|
595
|
+
self, status: str = "success", error: BaseException | None = None
|
|
596
|
+
) -> None:
|
|
597
|
+
response = self.last
|
|
598
|
+
if not error:
|
|
599
|
+
final = getattr(self.stream, "get_final_message", None)
|
|
600
|
+
if callable(final):
|
|
601
|
+
try:
|
|
602
|
+
response = final()
|
|
603
|
+
if inspect.isawaitable(response):
|
|
604
|
+
close = getattr(response, "close", None)
|
|
605
|
+
if close:
|
|
606
|
+
close()
|
|
607
|
+
response = self.last
|
|
608
|
+
except Exception:
|
|
609
|
+
pass
|
|
610
|
+
self.call.finish(
|
|
611
|
+
response,
|
|
612
|
+
status=status,
|
|
613
|
+
error=error,
|
|
614
|
+
stream=True,
|
|
615
|
+
ttft_ms=self.ttft_ms,
|
|
616
|
+
response_text="".join(self.parts) or None,
|
|
617
|
+
)
|
|
618
|
+
|
|
619
|
+
async def finish_async(
|
|
620
|
+
self, status: str = "success", error: BaseException | None = None
|
|
621
|
+
) -> None:
|
|
622
|
+
response = self.last
|
|
623
|
+
if not error:
|
|
624
|
+
final = getattr(self.stream, "get_final_message", None)
|
|
625
|
+
if callable(final):
|
|
626
|
+
try:
|
|
627
|
+
response = final()
|
|
628
|
+
if inspect.isawaitable(response):
|
|
629
|
+
response = await response
|
|
630
|
+
except Exception:
|
|
631
|
+
response = self.last
|
|
632
|
+
self.call.finish(
|
|
633
|
+
response,
|
|
634
|
+
status=status,
|
|
635
|
+
error=error,
|
|
636
|
+
stream=True,
|
|
637
|
+
ttft_ms=self.ttft_ms,
|
|
638
|
+
response_text="".join(self.parts) or None,
|
|
639
|
+
)
|
|
640
|
+
|
|
641
|
+
|
|
642
|
+
class SyncStream:
|
|
643
|
+
def __init__(self, stream: Any, call: CallState) -> None:
|
|
644
|
+
self._state = _StreamState(stream, call)
|
|
645
|
+
|
|
646
|
+
def __getattr__(self, name: str) -> Any:
|
|
647
|
+
return getattr(self._state.stream, name)
|
|
648
|
+
|
|
649
|
+
def __iter__(self) -> "SyncStream":
|
|
650
|
+
if self._state.iterator is None:
|
|
651
|
+
self._state.iterator = iter(self._state.stream)
|
|
652
|
+
return self
|
|
653
|
+
|
|
654
|
+
def __next__(self) -> Any:
|
|
655
|
+
if self._state.iterator is None:
|
|
656
|
+
self._state.iterator = iter(self._state.stream)
|
|
657
|
+
while True:
|
|
658
|
+
try:
|
|
659
|
+
value = self._state.chunk(next(self._state.iterator))
|
|
660
|
+
if _usage_only_chunk(value, self._state.call):
|
|
661
|
+
continue
|
|
662
|
+
return value
|
|
663
|
+
except StopIteration:
|
|
664
|
+
self._state.finish()
|
|
665
|
+
raise
|
|
666
|
+
except BaseException as exc:
|
|
667
|
+
self._state.finish(error=exc)
|
|
668
|
+
raise
|
|
669
|
+
|
|
670
|
+
def __enter__(self) -> "SyncStream":
|
|
671
|
+
enter = getattr(self._state.stream, "__enter__", None)
|
|
672
|
+
if enter:
|
|
673
|
+
enter()
|
|
674
|
+
return self
|
|
675
|
+
|
|
676
|
+
def __exit__(self, exc_type, exc, tb) -> Any:
|
|
677
|
+
if exc:
|
|
678
|
+
self._state.finish(error=exc)
|
|
679
|
+
elif not self._state.call.done:
|
|
680
|
+
self._state.finish(status="abandoned")
|
|
681
|
+
exit_fn = getattr(self._state.stream, "__exit__", None)
|
|
682
|
+
return exit_fn(exc_type, exc, tb) if exit_fn else None
|
|
683
|
+
|
|
684
|
+
def close(self) -> None:
|
|
685
|
+
close = getattr(self._state.stream, "close", None)
|
|
686
|
+
if close:
|
|
687
|
+
close()
|
|
688
|
+
if not self._state.call.done:
|
|
689
|
+
self._state.finish(status="abandoned")
|
|
690
|
+
|
|
691
|
+
def __del__(self) -> None:
|
|
692
|
+
try:
|
|
693
|
+
if not self._state.call.done:
|
|
694
|
+
self._state.finish(status="abandoned")
|
|
695
|
+
except Exception:
|
|
696
|
+
pass
|
|
697
|
+
|
|
698
|
+
|
|
699
|
+
class AsyncStream:
|
|
700
|
+
def __init__(self, stream: Any, call: CallState) -> None:
|
|
701
|
+
self._state = _StreamState(stream, call)
|
|
702
|
+
|
|
703
|
+
def __getattr__(self, name: str) -> Any:
|
|
704
|
+
return getattr(self._state.stream, name)
|
|
705
|
+
|
|
706
|
+
def __aiter__(self) -> "AsyncStream":
|
|
707
|
+
if self._state.iterator is None:
|
|
708
|
+
self._state.iterator = self._state.stream.__aiter__()
|
|
709
|
+
return self
|
|
710
|
+
|
|
711
|
+
async def __anext__(self) -> Any:
|
|
712
|
+
if self._state.iterator is None:
|
|
713
|
+
self._state.iterator = self._state.stream.__aiter__()
|
|
714
|
+
while True:
|
|
715
|
+
try:
|
|
716
|
+
value = self._state.chunk(await self._state.iterator.__anext__())
|
|
717
|
+
if _usage_only_chunk(value, self._state.call):
|
|
718
|
+
continue
|
|
719
|
+
return value
|
|
720
|
+
except StopAsyncIteration:
|
|
721
|
+
await self._state.finish_async()
|
|
722
|
+
raise
|
|
723
|
+
except BaseException as exc:
|
|
724
|
+
await self._state.finish_async(error=exc)
|
|
725
|
+
raise
|
|
726
|
+
|
|
727
|
+
async def __aenter__(self) -> "AsyncStream":
|
|
728
|
+
enter = getattr(self._state.stream, "__aenter__", None)
|
|
729
|
+
if enter:
|
|
730
|
+
await enter()
|
|
731
|
+
return self
|
|
732
|
+
|
|
733
|
+
async def __aexit__(self, exc_type, exc, tb) -> Any:
|
|
734
|
+
if exc:
|
|
735
|
+
await self._state.finish_async(error=exc)
|
|
736
|
+
elif not self._state.call.done:
|
|
737
|
+
await self._state.finish_async(status="abandoned")
|
|
738
|
+
exit_fn = getattr(self._state.stream, "__aexit__", None)
|
|
739
|
+
return await exit_fn(exc_type, exc, tb) if exit_fn else None
|
|
740
|
+
|
|
741
|
+
async def aclose(self) -> None:
|
|
742
|
+
close = getattr(self._state.stream, "aclose", None)
|
|
743
|
+
if close:
|
|
744
|
+
await close()
|
|
745
|
+
if not self._state.call.done:
|
|
746
|
+
await self._state.finish_async(status="abandoned")
|
|
747
|
+
|
|
748
|
+
def __del__(self) -> None:
|
|
749
|
+
try:
|
|
750
|
+
if not self._state.call.done:
|
|
751
|
+
self._state.call.finish(
|
|
752
|
+
self._state.last,
|
|
753
|
+
status="abandoned",
|
|
754
|
+
stream=True,
|
|
755
|
+
ttft_ms=self._state.ttft_ms,
|
|
756
|
+
response_text="".join(self._state.parts) or None,
|
|
757
|
+
)
|
|
758
|
+
except Exception:
|
|
759
|
+
pass
|
|
760
|
+
|
|
761
|
+
|
|
762
|
+
_runtime: Runtime | None = None
|
|
763
|
+
_seen_batch_items: set[str] = set()
|
|
764
|
+
|
|
765
|
+
|
|
766
|
+
def set_runtime(runtime: Runtime | None) -> None:
|
|
767
|
+
global _runtime
|
|
768
|
+
_runtime = runtime
|
|
769
|
+
|
|
770
|
+
|
|
771
|
+
def _request(args: tuple, kwargs: dict) -> dict[str, Any]:
|
|
772
|
+
request: dict[str, Any] = {}
|
|
773
|
+
if args and isinstance(args[0], Mapping):
|
|
774
|
+
request.update(args[0])
|
|
775
|
+
request.update(kwargs)
|
|
776
|
+
return request
|
|
777
|
+
|
|
778
|
+
|
|
779
|
+
def _mark_batch_item(key: str) -> bool:
|
|
780
|
+
if key in _seen_batch_items:
|
|
781
|
+
return False
|
|
782
|
+
if len(_seen_batch_items) >= 100_000:
|
|
783
|
+
_seen_batch_items.clear()
|
|
784
|
+
_seen_batch_items.add(key)
|
|
785
|
+
return True
|
|
786
|
+
|
|
787
|
+
|
|
788
|
+
def _capture_openai_batch_item(
|
|
789
|
+
runtime: Runtime,
|
|
790
|
+
item: Any,
|
|
791
|
+
*,
|
|
792
|
+
source_id: str,
|
|
793
|
+
context: CaptureContext,
|
|
794
|
+
) -> None:
|
|
795
|
+
response = _get(item, "response")
|
|
796
|
+
error = _get(item, "error")
|
|
797
|
+
custom_id = str(_get(item, "custom_id") or "")
|
|
798
|
+
item_id = str(_get(item, "id") or custom_id)
|
|
799
|
+
if response is None and error is None:
|
|
800
|
+
return
|
|
801
|
+
response_id = str(_get(response, "request_id") or item_id)
|
|
802
|
+
if not _mark_batch_item(f"openai:{source_id}:{item_id}:{response_id}"):
|
|
803
|
+
return
|
|
804
|
+
body = _get(response, "body") or {}
|
|
805
|
+
normalized = dict(body) if isinstance(body, Mapping) else body
|
|
806
|
+
if isinstance(normalized, dict) and response_id:
|
|
807
|
+
normalized = {**normalized, "_request_id": response_id}
|
|
808
|
+
request = {
|
|
809
|
+
"model": _get(body, "model"),
|
|
810
|
+
"batch": True,
|
|
811
|
+
"service_tier": "batch",
|
|
812
|
+
"batch_custom_id": custom_id or None,
|
|
813
|
+
"batch_item_id": item_id or None,
|
|
814
|
+
}
|
|
815
|
+
object_type = str(_get(body, "object") or "")
|
|
816
|
+
endpoint = "batch.responses" if object_type == "response" else "batch.chat.completions"
|
|
817
|
+
call = runtime.call_state("openai", endpoint, request, context=context)
|
|
818
|
+
status_code = _int(_get(response, "status_code"))
|
|
819
|
+
failed = error is not None or (status_code is not None and status_code >= 400)
|
|
820
|
+
call.finish(normalized, status="error" if failed else None)
|
|
821
|
+
|
|
822
|
+
|
|
823
|
+
def _capture_openai_batch_content(
|
|
824
|
+
runtime: Runtime,
|
|
825
|
+
result: Any,
|
|
826
|
+
*,
|
|
827
|
+
source_id: str,
|
|
828
|
+
context: CaptureContext,
|
|
829
|
+
) -> None:
|
|
830
|
+
try:
|
|
831
|
+
content = result if isinstance(result, (str, bytes, bytearray)) else _get(result, "content")
|
|
832
|
+
if isinstance(content, (bytes, bytearray)):
|
|
833
|
+
content = bytes(content).decode("utf-8")
|
|
834
|
+
if not isinstance(content, str):
|
|
835
|
+
return
|
|
836
|
+
for line in content.splitlines():
|
|
837
|
+
if not line.strip():
|
|
838
|
+
continue
|
|
839
|
+
try:
|
|
840
|
+
item = json.loads(line)
|
|
841
|
+
except (TypeError, json.JSONDecodeError):
|
|
842
|
+
continue
|
|
843
|
+
if isinstance(item, Mapping) and "custom_id" in item and (
|
|
844
|
+
"response" in item or "error" in item
|
|
845
|
+
):
|
|
846
|
+
_capture_openai_batch_item(
|
|
847
|
+
runtime, item, source_id=source_id, context=context
|
|
848
|
+
)
|
|
849
|
+
except Exception:
|
|
850
|
+
return
|
|
851
|
+
|
|
852
|
+
|
|
853
|
+
def _capture_anthropic_batch_item(
|
|
854
|
+
runtime: Runtime,
|
|
855
|
+
item: Any,
|
|
856
|
+
*,
|
|
857
|
+
batch_id: str,
|
|
858
|
+
context: CaptureContext,
|
|
859
|
+
) -> None:
|
|
860
|
+
result = _get(item, "result")
|
|
861
|
+
result_type = str(_get(result, "type") or "")
|
|
862
|
+
message = _get(result, "message")
|
|
863
|
+
custom_id = str(_get(item, "custom_id") or "")
|
|
864
|
+
key = f"anthropic:{batch_id}:{custom_id}:{result_type}"
|
|
865
|
+
if not result_type or not _mark_batch_item(key):
|
|
866
|
+
return
|
|
867
|
+
request = {
|
|
868
|
+
"model": _get(message, "model"),
|
|
869
|
+
"batch": True,
|
|
870
|
+
"service_tier": "batch",
|
|
871
|
+
"batch_custom_id": custom_id or None,
|
|
872
|
+
"batch_id": batch_id or None,
|
|
873
|
+
}
|
|
874
|
+
call = runtime.call_state("anthropic", "batch.messages", request, context=context)
|
|
875
|
+
call.finish(message or {}, status=None if result_type == "succeeded" else "error")
|
|
876
|
+
|
|
877
|
+
|
|
878
|
+
class _SyncBatchResults:
|
|
879
|
+
def __init__(
|
|
880
|
+
self, result: Any, runtime: Runtime, batch_id: str, context: CaptureContext
|
|
881
|
+
) -> None:
|
|
882
|
+
self._result = result
|
|
883
|
+
self._runtime = runtime
|
|
884
|
+
self._batch_id = batch_id
|
|
885
|
+
self._context = context
|
|
886
|
+
|
|
887
|
+
def __getattr__(self, name: str) -> Any:
|
|
888
|
+
return getattr(self._result, name)
|
|
889
|
+
|
|
890
|
+
def __iter__(self):
|
|
891
|
+
for item in self._result:
|
|
892
|
+
try:
|
|
893
|
+
_capture_anthropic_batch_item(
|
|
894
|
+
self._runtime,
|
|
895
|
+
item,
|
|
896
|
+
batch_id=self._batch_id,
|
|
897
|
+
context=self._context,
|
|
898
|
+
)
|
|
899
|
+
except Exception:
|
|
900
|
+
pass
|
|
901
|
+
yield item
|
|
902
|
+
|
|
903
|
+
|
|
904
|
+
class _AsyncBatchResults:
|
|
905
|
+
def __init__(
|
|
906
|
+
self, result: Any, runtime: Runtime, batch_id: str, context: CaptureContext
|
|
907
|
+
) -> None:
|
|
908
|
+
self._result = result
|
|
909
|
+
self._runtime = runtime
|
|
910
|
+
self._batch_id = batch_id
|
|
911
|
+
self._context = context
|
|
912
|
+
|
|
913
|
+
def __getattr__(self, name: str) -> Any:
|
|
914
|
+
return getattr(self._result, name)
|
|
915
|
+
|
|
916
|
+
def __aiter__(self):
|
|
917
|
+
return self
|
|
918
|
+
|
|
919
|
+
async def __anext__(self):
|
|
920
|
+
item = await self._result.__anext__()
|
|
921
|
+
try:
|
|
922
|
+
_capture_anthropic_batch_item(
|
|
923
|
+
self._runtime,
|
|
924
|
+
item,
|
|
925
|
+
batch_id=self._batch_id,
|
|
926
|
+
context=self._context,
|
|
927
|
+
)
|
|
928
|
+
except Exception:
|
|
929
|
+
pass
|
|
930
|
+
return item
|
|
931
|
+
|
|
932
|
+
|
|
933
|
+
def _wrap_anthropic_batch_results(
|
|
934
|
+
result: Any, runtime: Runtime, batch_id: str, context: CaptureContext
|
|
935
|
+
) -> Any:
|
|
936
|
+
if hasattr(result, "__aiter__"):
|
|
937
|
+
iterator = result.__aiter__()
|
|
938
|
+
return _AsyncBatchResults(iterator, runtime, batch_id, context)
|
|
939
|
+
if hasattr(result, "__iter__") and not isinstance(result, (str, bytes, bytearray)):
|
|
940
|
+
return _SyncBatchResults(result, runtime, batch_id, context)
|
|
941
|
+
return result
|
|
942
|
+
|
|
943
|
+
|
|
944
|
+
def _patch_openai_batch_content(owner: Any, method_name: str) -> bool:
|
|
945
|
+
original = getattr(owner, method_name, None)
|
|
946
|
+
if not callable(original):
|
|
947
|
+
return False
|
|
948
|
+
if getattr(original, "__metergraph_batch__", False):
|
|
949
|
+
return True
|
|
950
|
+
|
|
951
|
+
@functools.wraps(original)
|
|
952
|
+
def wrapped(*args, **kwargs):
|
|
953
|
+
runtime = _runtime
|
|
954
|
+
if runtime is None:
|
|
955
|
+
return original(*args, **kwargs)
|
|
956
|
+
source_id = str(args[0] if args else kwargs.get("file_id") or "unknown")
|
|
957
|
+
context = snapshot()
|
|
958
|
+
result = original(*args, **kwargs)
|
|
959
|
+
if inspect.isawaitable(result):
|
|
960
|
+
|
|
961
|
+
async def await_result():
|
|
962
|
+
resolved = await result
|
|
963
|
+
_capture_openai_batch_content(
|
|
964
|
+
runtime, resolved, source_id=source_id, context=context
|
|
965
|
+
)
|
|
966
|
+
return resolved
|
|
967
|
+
|
|
968
|
+
return await_result()
|
|
969
|
+
_capture_openai_batch_content(
|
|
970
|
+
runtime, result, source_id=source_id, context=context
|
|
971
|
+
)
|
|
972
|
+
return result
|
|
973
|
+
|
|
974
|
+
wrapped.__metergraph_batch__ = True # type: ignore[attr-defined]
|
|
975
|
+
try:
|
|
976
|
+
setattr(owner, method_name, wrapped)
|
|
977
|
+
except Exception:
|
|
978
|
+
return False
|
|
979
|
+
return True
|
|
980
|
+
|
|
981
|
+
|
|
982
|
+
def _patch_anthropic_batch_results(owner: Any) -> bool:
|
|
983
|
+
original = getattr(owner, "results", None)
|
|
984
|
+
if not callable(original):
|
|
985
|
+
return False
|
|
986
|
+
if getattr(original, "__metergraph_batch__", False):
|
|
987
|
+
return True
|
|
988
|
+
|
|
989
|
+
@functools.wraps(original)
|
|
990
|
+
def wrapped(*args, **kwargs):
|
|
991
|
+
runtime = _runtime
|
|
992
|
+
if runtime is None:
|
|
993
|
+
return original(*args, **kwargs)
|
|
994
|
+
batch_id = str(
|
|
995
|
+
args[0]
|
|
996
|
+
if args
|
|
997
|
+
else kwargs.get("message_batch_id") or kwargs.get("batch_id") or "unknown"
|
|
998
|
+
)
|
|
999
|
+
context = snapshot()
|
|
1000
|
+
result = original(*args, **kwargs)
|
|
1001
|
+
if inspect.isawaitable(result):
|
|
1002
|
+
|
|
1003
|
+
async def await_result():
|
|
1004
|
+
resolved = await result
|
|
1005
|
+
return _wrap_anthropic_batch_results(
|
|
1006
|
+
resolved, runtime, batch_id, context
|
|
1007
|
+
)
|
|
1008
|
+
|
|
1009
|
+
return await_result()
|
|
1010
|
+
return _wrap_anthropic_batch_results(result, runtime, batch_id, context)
|
|
1011
|
+
|
|
1012
|
+
wrapped.__metergraph_batch__ = True # type: ignore[attr-defined]
|
|
1013
|
+
try:
|
|
1014
|
+
setattr(owner, "results", wrapped)
|
|
1015
|
+
except Exception:
|
|
1016
|
+
return False
|
|
1017
|
+
return True
|
|
1018
|
+
|
|
1019
|
+
|
|
1020
|
+
def _patch(owner: Any, method_name: str, provider: str, endpoint: str) -> bool:
|
|
1021
|
+
original = getattr(owner, method_name, None)
|
|
1022
|
+
if not callable(original):
|
|
1023
|
+
return False
|
|
1024
|
+
if getattr(original, "__metergraph__", False):
|
|
1025
|
+
return True
|
|
1026
|
+
|
|
1027
|
+
@functools.wraps(original)
|
|
1028
|
+
def wrapped(*args, **kwargs):
|
|
1029
|
+
runtime = _runtime
|
|
1030
|
+
if runtime is None:
|
|
1031
|
+
return original(*args, **kwargs)
|
|
1032
|
+
if (
|
|
1033
|
+
provider == "openai"
|
|
1034
|
+
and endpoint == "chat.completions"
|
|
1035
|
+
and os.getenv("METERGRAPH_PATCH_STREAM_USAGE", "1") != "0"
|
|
1036
|
+
):
|
|
1037
|
+
incoming = _request(args, kwargs)
|
|
1038
|
+
if incoming.get("stream") is True and "stream_options" not in incoming:
|
|
1039
|
+
if args and isinstance(args[0], Mapping):
|
|
1040
|
+
first = {**args[0], "stream_options": {"include_usage": True}}
|
|
1041
|
+
args = (first, *args[1:])
|
|
1042
|
+
else:
|
|
1043
|
+
kwargs = {**kwargs, "stream_options": {"include_usage": True}}
|
|
1044
|
+
request = _request(args, kwargs)
|
|
1045
|
+
call = runtime.call_state(provider, endpoint, request)
|
|
1046
|
+
try:
|
|
1047
|
+
result = original(*args, **kwargs)
|
|
1048
|
+
except BaseException as exc:
|
|
1049
|
+
call.finish(error=exc)
|
|
1050
|
+
raise
|
|
1051
|
+
|
|
1052
|
+
if inspect.isawaitable(result):
|
|
1053
|
+
|
|
1054
|
+
async def await_result():
|
|
1055
|
+
try:
|
|
1056
|
+
resolved = await result
|
|
1057
|
+
except BaseException as exc:
|
|
1058
|
+
call.finish(error=exc)
|
|
1059
|
+
raise
|
|
1060
|
+
return _finish_or_stream(resolved, call, endpoint, request)
|
|
1061
|
+
|
|
1062
|
+
return await_result()
|
|
1063
|
+
return _finish_or_stream(result, call, endpoint, request)
|
|
1064
|
+
|
|
1065
|
+
wrapped.__metergraph__ = True # type: ignore[attr-defined]
|
|
1066
|
+
try:
|
|
1067
|
+
setattr(owner, method_name, wrapped)
|
|
1068
|
+
except Exception:
|
|
1069
|
+
return False
|
|
1070
|
+
return True
|
|
1071
|
+
|
|
1072
|
+
|
|
1073
|
+
def _finish_or_stream(
|
|
1074
|
+
result: Any, call: CallState, endpoint: str, request: Mapping[str, Any]
|
|
1075
|
+
):
|
|
1076
|
+
is_stream = endpoint.endswith(".stream") or bool(request.get("stream"))
|
|
1077
|
+
if is_stream and hasattr(result, "__aiter__"):
|
|
1078
|
+
return AsyncStream(result, call)
|
|
1079
|
+
if is_stream and hasattr(result, "__iter__"):
|
|
1080
|
+
return SyncStream(result, call)
|
|
1081
|
+
call.finish(result)
|
|
1082
|
+
return result
|
|
1083
|
+
|
|
1084
|
+
|
|
1085
|
+
def wrap(client: Any, *, provider: str | None = None) -> Any:
|
|
1086
|
+
"""Patch supported resource methods on an OpenAI, Anthropic, or Google client."""
|
|
1087
|
+
if provider is None:
|
|
1088
|
+
if hasattr(getattr(client, "models", None), "generate_content"):
|
|
1089
|
+
provider = "google"
|
|
1090
|
+
elif hasattr(client, "chat") or hasattr(client, "responses"):
|
|
1091
|
+
provider = "openai"
|
|
1092
|
+
else:
|
|
1093
|
+
provider = "anthropic"
|
|
1094
|
+
seams: list[tuple[Any, str, str]] = []
|
|
1095
|
+
if provider == "google":
|
|
1096
|
+
for models in (
|
|
1097
|
+
getattr(client, "models", None),
|
|
1098
|
+
getattr(getattr(client, "aio", None), "models", None),
|
|
1099
|
+
):
|
|
1100
|
+
if models is not None:
|
|
1101
|
+
seams.extend(
|
|
1102
|
+
(
|
|
1103
|
+
(models, "generate_content", "models.generate_content"),
|
|
1104
|
+
(
|
|
1105
|
+
models,
|
|
1106
|
+
"generate_content_stream",
|
|
1107
|
+
"models.generate_content.stream",
|
|
1108
|
+
),
|
|
1109
|
+
)
|
|
1110
|
+
)
|
|
1111
|
+
chat = getattr(getattr(client, "chat", None), "completions", None)
|
|
1112
|
+
if chat is not None:
|
|
1113
|
+
seams.append((chat, "create", "chat.completions"))
|
|
1114
|
+
responses = getattr(client, "responses", None)
|
|
1115
|
+
if responses is not None:
|
|
1116
|
+
seams.extend(
|
|
1117
|
+
(
|
|
1118
|
+
(responses, "create", "responses"),
|
|
1119
|
+
(responses, "stream", "responses.stream"),
|
|
1120
|
+
)
|
|
1121
|
+
)
|
|
1122
|
+
messages = getattr(client, "messages", None)
|
|
1123
|
+
if messages is not None:
|
|
1124
|
+
seams.extend(
|
|
1125
|
+
((messages, "create", "messages"), (messages, "stream", "messages.stream"))
|
|
1126
|
+
)
|
|
1127
|
+
patched = sum(
|
|
1128
|
+
_patch(owner, method, provider, endpoint) for owner, method, endpoint in seams
|
|
1129
|
+
)
|
|
1130
|
+
if provider == "openai":
|
|
1131
|
+
files = getattr(client, "files", None)
|
|
1132
|
+
if files is not None:
|
|
1133
|
+
patched += int(_patch_openai_batch_content(files, "content"))
|
|
1134
|
+
patched += int(_patch_openai_batch_content(files, "retrieve_content"))
|
|
1135
|
+
elif provider == "anthropic":
|
|
1136
|
+
batch_owners = [getattr(messages, "batches", None)]
|
|
1137
|
+
beta_messages = getattr(getattr(client, "beta", None), "messages", None)
|
|
1138
|
+
batch_owners.append(getattr(beta_messages, "batches", None))
|
|
1139
|
+
patched += sum(
|
|
1140
|
+
int(_patch_anthropic_batch_results(owner))
|
|
1141
|
+
for owner in batch_owners
|
|
1142
|
+
if owner is not None
|
|
1143
|
+
)
|
|
1144
|
+
if not patched:
|
|
1145
|
+
log.warning("Metergraph found no supported methods on %s client", provider)
|
|
1146
|
+
return client
|