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/_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