tunarag-python 0.2.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,270 @@
1
+ """Adapters for LangChain runnables and compiled LangGraph graphs."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import inspect
6
+ import time
7
+ from collections.abc import Callable, Mapping, Sequence
8
+ from typing import Any, Protocol
9
+
10
+ from ..dataset import EvaluationExample
11
+ from ..domain import Candidate, UsageRecord
12
+
13
+
14
+ class AsyncRunnable(Protocol):
15
+ """Shared asynchronous invocation surface implemented by both frameworks."""
16
+
17
+ async def ainvoke(
18
+ self,
19
+ input: Any,
20
+ config: Any = None,
21
+ **kwargs: Any,
22
+ ) -> Any: ...
23
+
24
+
25
+ RunnableFactory = Callable[[Candidate], Any]
26
+ RunnableInputMapper = Callable[[Any, Candidate], Any]
27
+ RunnableConfigMapper = Callable[[Candidate], Mapping[str, Any] | None]
28
+ RunnableOutputMapper = Callable[[Any], Any]
29
+ RunnableUsageMapper = Callable[[Any], Sequence[UsageRecord]]
30
+
31
+
32
+ class _RunnableAdapter:
33
+ def __init__(
34
+ self,
35
+ runnable: Any,
36
+ *,
37
+ factory: RunnableFactory | None,
38
+ input_key: str,
39
+ input_mapper: RunnableInputMapper | None,
40
+ config_mapper: RunnableConfigMapper | None,
41
+ output_mapper: RunnableOutputMapper | None,
42
+ usage_mapper: RunnableUsageMapper | None,
43
+ capture_message_usage: bool,
44
+ identity: Any,
45
+ framework: str,
46
+ ) -> None:
47
+ if (runnable is None) == (factory is None):
48
+ raise ValueError("provide exactly one of runnable or factory")
49
+ if runnable is not None:
50
+ _require_runnable(runnable)
51
+ if factory is not None and not callable(factory):
52
+ raise TypeError("runnable factory must be callable")
53
+ if not input_key.strip():
54
+ raise ValueError("runnable input key must not be empty")
55
+ for name, mapper in (
56
+ ("input_mapper", input_mapper),
57
+ ("config_mapper", config_mapper),
58
+ ("output_mapper", output_mapper),
59
+ ("usage_mapper", usage_mapper),
60
+ ):
61
+ if mapper is not None and not callable(mapper):
62
+ raise TypeError(f"{name} must be callable")
63
+ if not isinstance(capture_message_usage, bool):
64
+ raise TypeError("capture_message_usage must be a boolean")
65
+ self._runnable = runnable
66
+ self._factory = factory
67
+ self._input_key = input_key
68
+ self._input_mapper = input_mapper
69
+ self._config_mapper = config_mapper
70
+ self._output_mapper = output_mapper
71
+ self._usage_mapper = usage_mapper
72
+ self._capture_message_usage = capture_message_usage
73
+ self._identity = identity
74
+ self._framework = framework
75
+
76
+ async def run(self, candidate: Candidate, example: Any) -> tuple[Any, tuple[UsageRecord, ...]]:
77
+ """Invoke one framework runnable for a TunaRAG candidate and example."""
78
+
79
+ runnable = self._resolve_runnable(candidate)
80
+ runnable_input = self._map_input(example, candidate)
81
+ config = self._map_config(candidate)
82
+ started_at = time.perf_counter()
83
+ raw_output = await runnable.ainvoke(runnable_input, config=config)
84
+ latency = time.perf_counter() - started_at
85
+ output = raw_output if self._output_mapper is None else self._output_mapper(raw_output)
86
+ usage = list(self._map_usage(raw_output))
87
+ usage.append(
88
+ UsageRecord(
89
+ component=f"{self._framework}.runnable",
90
+ latency_seconds=latency,
91
+ )
92
+ )
93
+ return output, tuple(usage)
94
+
95
+ def fingerprint(self) -> dict[str, Any]:
96
+ """Return stable, provider-secret-free adapter semantics."""
97
+
98
+ runnable_identity = _type_name(self._runnable) if self._runnable is not None else None
99
+ return {
100
+ "integration": self._framework,
101
+ "schema_version": 1,
102
+ "runnable": runnable_identity,
103
+ "factory": _callable_identity(self._factory),
104
+ "input_key": self._input_key,
105
+ "input_mapper": _callable_identity(self._input_mapper),
106
+ "config_mapper": _callable_identity(self._config_mapper),
107
+ "output_mapper": _callable_identity(self._output_mapper),
108
+ "usage_mapper": _callable_identity(self._usage_mapper),
109
+ "capture_message_usage": self._capture_message_usage,
110
+ "identity": self._identity,
111
+ }
112
+
113
+ def _resolve_runnable(self, candidate: Candidate) -> Any:
114
+ if self._runnable is not None:
115
+ return self._runnable
116
+ if self._factory is None:
117
+ raise RuntimeError("runnable adapter is not configured")
118
+ runnable = self._factory(candidate)
119
+ _require_runnable(runnable)
120
+ return runnable
121
+
122
+ def _map_input(self, example: Any, candidate: Candidate) -> Any:
123
+ if self._input_mapper is not None:
124
+ return self._input_mapper(example, candidate)
125
+ if isinstance(example, EvaluationExample):
126
+ return {self._input_key: example.query}
127
+ return example
128
+
129
+ def _map_config(self, candidate: Candidate) -> Mapping[str, Any] | None:
130
+ if self._config_mapper is None:
131
+ return {"configurable": dict(candidate.parameters)}
132
+ config = self._config_mapper(candidate)
133
+ if config is not None and not isinstance(config, Mapping):
134
+ raise TypeError("config_mapper must return a mapping or None")
135
+ return config
136
+
137
+ def _map_usage(self, raw_output: Any) -> tuple[UsageRecord, ...]:
138
+ if self._usage_mapper is not None:
139
+ usage = tuple(self._usage_mapper(raw_output))
140
+ if any(not isinstance(item, UsageRecord) for item in usage):
141
+ raise TypeError("usage_mapper must return only UsageRecord values")
142
+ return usage
143
+ if not self._capture_message_usage:
144
+ return ()
145
+ return _message_usage(raw_output, framework=self._framework)
146
+
147
+
148
+ class LangChainAdapter(_RunnableAdapter):
149
+ """Expose a LangChain runnable through TunaRAG's standard adapter contract."""
150
+
151
+ def __init__(
152
+ self,
153
+ runnable: Any = None,
154
+ *,
155
+ factory: RunnableFactory | None = None,
156
+ input_key: str = "question",
157
+ input_mapper: RunnableInputMapper | None = None,
158
+ config_mapper: RunnableConfigMapper | None = None,
159
+ output_mapper: RunnableOutputMapper | None = None,
160
+ usage_mapper: RunnableUsageMapper | None = None,
161
+ capture_message_usage: bool = True,
162
+ identity: Any = None,
163
+ ) -> None:
164
+ super().__init__(
165
+ runnable,
166
+ factory=factory,
167
+ input_key=input_key,
168
+ input_mapper=input_mapper,
169
+ config_mapper=config_mapper,
170
+ output_mapper=output_mapper,
171
+ usage_mapper=usage_mapper,
172
+ capture_message_usage=capture_message_usage,
173
+ identity=identity,
174
+ framework="langchain",
175
+ )
176
+
177
+
178
+ class LangGraphAdapter(_RunnableAdapter):
179
+ """Expose a compiled LangGraph graph through TunaRAG's adapter contract."""
180
+
181
+ def __init__(
182
+ self,
183
+ graph: Any = None,
184
+ *,
185
+ factory: RunnableFactory | None = None,
186
+ input_key: str = "question",
187
+ input_mapper: RunnableInputMapper | None = None,
188
+ config_mapper: RunnableConfigMapper | None = None,
189
+ output_mapper: RunnableOutputMapper | None = None,
190
+ usage_mapper: RunnableUsageMapper | None = None,
191
+ capture_message_usage: bool = True,
192
+ identity: Any = None,
193
+ ) -> None:
194
+ super().__init__(
195
+ graph,
196
+ factory=factory,
197
+ input_key=input_key,
198
+ input_mapper=input_mapper,
199
+ config_mapper=config_mapper,
200
+ output_mapper=output_mapper,
201
+ usage_mapper=usage_mapper,
202
+ capture_message_usage=capture_message_usage,
203
+ identity=identity,
204
+ framework="langgraph",
205
+ )
206
+
207
+
208
+ def _message_usage(output: Any, *, framework: str) -> tuple[UsageRecord, ...]:
209
+ records: list[UsageRecord] = []
210
+ visited: set[int] = set()
211
+
212
+ def collect(value: Any) -> None:
213
+ if value is None or isinstance(value, (str, bytes, bytearray, bool, int, float)):
214
+ return
215
+ value_id = id(value)
216
+ if value_id in visited:
217
+ return
218
+ visited.add(value_id)
219
+ metadata = (
220
+ value.get("usage_metadata")
221
+ if isinstance(value, Mapping)
222
+ else getattr(value, "usage_metadata", None)
223
+ )
224
+ if isinstance(metadata, Mapping):
225
+ input_tokens = _token_count(metadata.get("input_tokens"))
226
+ output_tokens = _token_count(metadata.get("output_tokens"))
227
+ if input_tokens is not None or output_tokens is not None:
228
+ records.append(
229
+ UsageRecord(
230
+ component=f"{framework}.message",
231
+ input_tokens=input_tokens,
232
+ output_tokens=output_tokens,
233
+ )
234
+ )
235
+ if isinstance(value, Mapping):
236
+ for item in value.values():
237
+ collect(item)
238
+ elif isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
239
+ for item in value:
240
+ collect(item)
241
+
242
+ collect(output)
243
+ return tuple(records)
244
+
245
+
246
+ def _token_count(value: Any) -> int | None:
247
+ if value is None:
248
+ return None
249
+ if isinstance(value, bool) or not isinstance(value, int) or value < 0:
250
+ return None
251
+ return value
252
+
253
+
254
+ def _require_runnable(value: Any) -> None:
255
+ method = getattr(value, "ainvoke", None)
256
+ if method is None or not callable(method):
257
+ raise TypeError("framework object must provide an async ainvoke method")
258
+
259
+
260
+ def _callable_identity(value: Callable[..., Any] | None) -> str | None:
261
+ if value is None:
262
+ return None
263
+ if inspect.ismethod(value) or inspect.isfunction(value):
264
+ return f"{value.__module__}.{value.__qualname__}"
265
+ return _type_name(value)
266
+
267
+
268
+ def _type_name(value: Any) -> str:
269
+ value_type = type(value)
270
+ return f"{value_type.__module__}.{value_type.__qualname__}"
tunarag/objective.py ADDED
@@ -0,0 +1,75 @@
1
+ """Objective aggregation for quality, cost, and latency metrics."""
2
+
3
+ from collections.abc import Sequence
4
+ from dataclasses import dataclass
5
+ from enum import Enum
6
+
7
+ from .domain import MetricValue, ValueStatus
8
+
9
+
10
+ class ObjectiveDirection(str, Enum):
11
+ """Whether larger or smaller metric values are beneficial."""
12
+
13
+ MAXIMIZE = "maximize"
14
+ MINIMIZE = "minimize"
15
+
16
+
17
+ @dataclass(frozen=True, slots=True)
18
+ class ObjectiveTerm:
19
+ """One weighted metric contribution to an objective score."""
20
+
21
+ metric: str
22
+ weight: float = 1.0
23
+ direction: ObjectiveDirection = ObjectiveDirection.MAXIMIZE
24
+ minimum: float | None = None
25
+ maximum: float | None = None
26
+
27
+ def __post_init__(self) -> None:
28
+ if not self.metric.strip():
29
+ raise ValueError("objective metric must not be empty")
30
+ if self.weight < 0:
31
+ raise ValueError("objective weight must not be negative")
32
+ if self.minimum is not None and self.maximum is not None and self.minimum >= self.maximum:
33
+ raise ValueError("objective minimum must be less than maximum")
34
+ if (self.minimum is None) != (self.maximum is None):
35
+ raise ValueError("objective minimum and maximum must be provided together")
36
+
37
+
38
+ @dataclass(frozen=True, slots=True)
39
+ class ObjectiveResult:
40
+ """Aggregated score and confidence state."""
41
+
42
+ score: float | None
43
+ status: ValueStatus
44
+
45
+
46
+ @dataclass(frozen=True, slots=True)
47
+ class Objective:
48
+ """Combine metric observations into a benefit-oriented weighted score."""
49
+
50
+ terms: tuple[ObjectiveTerm, ...]
51
+
52
+ def __post_init__(self) -> None:
53
+ if not self.terms:
54
+ raise ValueError("objective must contain at least one term")
55
+ if sum(term.weight for term in self.terms) == 0:
56
+ raise ValueError("objective must contain a positive weight")
57
+
58
+ def evaluate(self, metrics: Sequence[MetricValue]) -> ObjectiveResult:
59
+ by_name = {metric.name: metric for metric in metrics}
60
+ score = 0.0
61
+ status = ValueStatus.EXACT
62
+ for term in self.terms:
63
+ metric = by_name.get(term.metric)
64
+ if metric is None or metric.value is None:
65
+ return ObjectiveResult(None, ValueStatus.UNAVAILABLE)
66
+ value = float(metric.value)
67
+ if term.minimum is not None and term.maximum is not None:
68
+ value = (value - term.minimum) / (term.maximum - term.minimum)
69
+ value = max(0.0, min(1.0, value))
70
+ if term.direction is ObjectiveDirection.MINIMIZE:
71
+ value = 1.0 - value if term.minimum is not None else -value
72
+ score += term.weight * value
73
+ if metric.status is ValueStatus.ESTIMATED:
74
+ status = ValueStatus.ESTIMATED
75
+ return ObjectiveResult(score, status)
tunarag/py.typed ADDED
@@ -0,0 +1 @@
1
+
tunarag/result.py ADDED
@@ -0,0 +1,309 @@
1
+ """Inspectable optimization result models."""
2
+
3
+ import csv
4
+ import json
5
+ from collections.abc import Mapping
6
+ from dataclasses import dataclass
7
+ from pathlib import Path
8
+ from typing import Any, cast
9
+
10
+ from .domain import Candidate, MetricValue, UsageRecord, ValueStatus
11
+ from .errors import redact
12
+ from .objective import ObjectiveDirection, ObjectiveResult
13
+ from .store import StudyStatus, TrialStatus
14
+
15
+
16
+ @dataclass(frozen=True, slots=True)
17
+ class UsageSummary:
18
+ """Aggregate usage with estimate provenance retained."""
19
+
20
+ input_tokens: int
21
+ output_tokens: int
22
+ exact_cost: float
23
+ estimated_cost: float
24
+ latency_seconds: float
25
+ status: ValueStatus
26
+
27
+
28
+ @dataclass(frozen=True, slots=True)
29
+ class TrialSummary:
30
+ """Actionable terminal summary for one candidate trial."""
31
+
32
+ id: str
33
+ sequence: int
34
+ candidate: Candidate
35
+ status: TrialStatus
36
+ metrics: tuple[MetricValue, ...]
37
+ usage: tuple[UsageRecord, ...]
38
+ objective: ObjectiveResult | None
39
+ attempts: int
40
+ cached: bool = False
41
+ error: str | None = None
42
+
43
+ def metric(self, name: str) -> MetricValue | None:
44
+ """Return one named metric, or None when it was not observed."""
45
+
46
+ return next((metric for metric in self.metrics if metric.name == name), None)
47
+
48
+
49
+ @dataclass(frozen=True, slots=True)
50
+ class OptimizationResult:
51
+ """Completed study result, including failures and no-success outcomes."""
52
+
53
+ study_id: str
54
+ study_status: StudyStatus
55
+ stop_reason: str
56
+ trials: tuple[TrialSummary, ...]
57
+ best_trial: TrialSummary | None
58
+ usage: UsageSummary
59
+
60
+ @property
61
+ def best_candidate(self) -> Candidate | None:
62
+ """Return the highest-scoring candidate when one succeeded."""
63
+
64
+ return None if self.best_trial is None else self.best_trial.candidate
65
+
66
+ @property
67
+ def successful_trials(self) -> tuple[TrialSummary, ...]:
68
+ """Return successful trials in execution order."""
69
+
70
+ return tuple(trial for trial in self.trials if trial.status is TrialStatus.SUCCEEDED)
71
+
72
+ @property
73
+ def failed_trials(self) -> tuple[TrialSummary, ...]:
74
+ """Return failed trials in execution order."""
75
+
76
+ return tuple(trial for trial in self.trials if trial.status is TrialStatus.FAILED)
77
+
78
+ @property
79
+ def ranked_trials(self) -> tuple[TrialSummary, ...]:
80
+ """Return successful trials by descending objective, then sequence."""
81
+
82
+ return tuple(
83
+ sorted(
84
+ (
85
+ trial
86
+ for trial in self.successful_trials
87
+ if trial.objective is not None and trial.objective.score is not None
88
+ ),
89
+ key=lambda trial: (-_objective_score(trial), trial.sequence),
90
+ )
91
+ )
92
+
93
+ @property
94
+ def metric_names(self) -> tuple[str, ...]:
95
+ """Return all observed metric names in lexical order."""
96
+
97
+ return tuple(sorted({metric.name for trial in self.trials for metric in trial.metrics}))
98
+
99
+ def metric_leaders(
100
+ self,
101
+ name: str,
102
+ *,
103
+ direction: ObjectiveDirection = ObjectiveDirection.MAXIMIZE,
104
+ ) -> tuple[TrialSummary, ...]:
105
+ """Rank successful trials that have an available named metric."""
106
+
107
+ if not name.strip():
108
+ raise ValueError("metric name must not be empty")
109
+ trials = [
110
+ trial
111
+ for trial in self.successful_trials
112
+ if (metric := trial.metric(name)) is not None and metric.value is not None
113
+ ]
114
+ multiplier = -1.0 if direction is ObjectiveDirection.MAXIMIZE else 1.0
115
+ return tuple(
116
+ sorted(
117
+ trials,
118
+ key=lambda trial: (
119
+ multiplier * cast(float, cast(MetricValue, trial.metric(name)).value),
120
+ trial.sequence,
121
+ ),
122
+ )
123
+ )
124
+
125
+ def as_dict(self) -> dict[str, Any]:
126
+ """Return a redacted, JSON-compatible report."""
127
+
128
+ return {
129
+ "schema_version": 1,
130
+ "study_id": self.study_id,
131
+ "study_status": self.study_status.value,
132
+ "stop_reason": self.stop_reason,
133
+ "best_trial_id": None if self.best_trial is None else self.best_trial.id,
134
+ "counts": {
135
+ "trials": len(self.trials),
136
+ "successful": len(self.successful_trials),
137
+ "failed": len(self.failed_trials),
138
+ },
139
+ "usage": _usage_dict(self.usage),
140
+ "trials": [_trial_dict(trial) for trial in self.trials],
141
+ }
142
+
143
+ def to_json(self, *, indent: int | None = 2) -> str:
144
+ """Serialize the redacted report to deterministic JSON."""
145
+
146
+ return json.dumps(
147
+ self.as_dict(),
148
+ ensure_ascii=False,
149
+ allow_nan=False,
150
+ indent=indent,
151
+ sort_keys=True,
152
+ )
153
+
154
+ def write_json(self, path: str | Path, *, indent: int | None = 2) -> Path:
155
+ """Write a redacted JSON report and return its path."""
156
+
157
+ target = Path(path)
158
+ target.write_text(self.to_json(indent=indent) + "\n", encoding="utf-8")
159
+ return target
160
+
161
+ def write_csv(self, path: str | Path) -> Path:
162
+ """Write one redacted, flattened row per trial and return its path."""
163
+
164
+ target = Path(path)
165
+ metric_columns = [
166
+ column
167
+ for name in self.metric_names
168
+ for column in (
169
+ f"metric.{name}.value",
170
+ f"metric.{name}.status",
171
+ f"metric.{name}.coverage",
172
+ )
173
+ ]
174
+ fieldnames = [
175
+ "study_id",
176
+ "trial_id",
177
+ "sequence",
178
+ "status",
179
+ "objective_score",
180
+ "objective_status",
181
+ "attempts",
182
+ "cached",
183
+ "error",
184
+ "candidate",
185
+ "input_tokens",
186
+ "output_tokens",
187
+ "exact_cost",
188
+ "estimated_cost",
189
+ "latency_seconds",
190
+ "usage_status",
191
+ *metric_columns,
192
+ ]
193
+ with target.open("w", encoding="utf-8", newline="") as stream:
194
+ writer = csv.DictWriter(stream, fieldnames=fieldnames, extrasaction="raise")
195
+ writer.writeheader()
196
+ for trial in self.trials:
197
+ writer.writerow(_trial_csv_row(self.study_id, trial, self.metric_names))
198
+ return target
199
+
200
+
201
+ def _usage_dict(usage: UsageSummary) -> dict[str, Any]:
202
+ return {
203
+ "input_tokens": usage.input_tokens,
204
+ "output_tokens": usage.output_tokens,
205
+ "exact_cost": usage.exact_cost,
206
+ "estimated_cost": usage.estimated_cost,
207
+ "latency_seconds": usage.latency_seconds,
208
+ "status": usage.status.value,
209
+ }
210
+
211
+
212
+ def _trial_dict(trial: TrialSummary) -> dict[str, Any]:
213
+ candidate = cast(Mapping[str, Any], redact(trial.candidate.parameters))
214
+ return {
215
+ "id": trial.id,
216
+ "sequence": trial.sequence,
217
+ "status": trial.status.value,
218
+ "candidate": dict(candidate),
219
+ "metrics": [
220
+ {
221
+ "name": metric.name,
222
+ "value": metric.value,
223
+ "status": metric.status.value,
224
+ "coverage": metric.coverage,
225
+ }
226
+ for metric in trial.metrics
227
+ ],
228
+ "usage": [
229
+ {
230
+ "component": item.component,
231
+ "input_tokens": item.input_tokens,
232
+ "output_tokens": item.output_tokens,
233
+ "cost": item.cost,
234
+ "latency_seconds": item.latency_seconds,
235
+ "status": item.status.value,
236
+ "pricing_version": item.pricing_version,
237
+ }
238
+ for item in trial.usage
239
+ ],
240
+ "objective": (
241
+ None
242
+ if trial.objective is None
243
+ else {"score": trial.objective.score, "status": trial.objective.status.value}
244
+ ),
245
+ "attempts": trial.attempts,
246
+ "cached": trial.cached,
247
+ "error": trial.error,
248
+ }
249
+
250
+
251
+ def _trial_csv_row(
252
+ study_id: str, trial: TrialSummary, metric_names: tuple[str, ...]
253
+ ) -> dict[str, Any]:
254
+ usage = _summarize_usage(trial.usage)
255
+ candidate = redact(trial.candidate.parameters)
256
+ row: dict[str, Any] = {
257
+ "study_id": study_id,
258
+ "trial_id": trial.id,
259
+ "sequence": trial.sequence,
260
+ "status": trial.status.value,
261
+ "objective_score": None if trial.objective is None else trial.objective.score,
262
+ "objective_status": None if trial.objective is None else trial.objective.status.value,
263
+ "attempts": trial.attempts,
264
+ "cached": trial.cached,
265
+ "error": trial.error,
266
+ "candidate": json.dumps(
267
+ candidate, ensure_ascii=False, sort_keys=True, separators=(",", ":")
268
+ ),
269
+ "input_tokens": usage.input_tokens,
270
+ "output_tokens": usage.output_tokens,
271
+ "exact_cost": usage.exact_cost,
272
+ "estimated_cost": usage.estimated_cost,
273
+ "latency_seconds": usage.latency_seconds,
274
+ "usage_status": usage.status.value,
275
+ }
276
+ for name in metric_names:
277
+ metric = trial.metric(name)
278
+ row[f"metric.{name}.value"] = None if metric is None else metric.value
279
+ row[f"metric.{name}.status"] = None if metric is None else metric.status.value
280
+ row[f"metric.{name}.coverage"] = None if metric is None else metric.coverage
281
+ return row
282
+
283
+
284
+ def _objective_score(trial: TrialSummary) -> float:
285
+ if trial.objective is None or trial.objective.score is None:
286
+ raise ValueError("ranked trial has no objective score")
287
+ return trial.objective.score
288
+
289
+
290
+ def _summarize_usage(records: tuple[UsageRecord, ...]) -> UsageSummary:
291
+ statuses = {record.status for record in records}
292
+ if not records or statuses == {ValueStatus.UNAVAILABLE}:
293
+ status = ValueStatus.UNAVAILABLE
294
+ elif statuses == {ValueStatus.EXACT}:
295
+ status = ValueStatus.EXACT
296
+ else:
297
+ status = ValueStatus.ESTIMATED
298
+ return UsageSummary(
299
+ input_tokens=sum(record.input_tokens or 0 for record in records),
300
+ output_tokens=sum(record.output_tokens or 0 for record in records),
301
+ exact_cost=sum(
302
+ record.cost or 0.0 for record in records if record.status is ValueStatus.EXACT
303
+ ),
304
+ estimated_cost=sum(
305
+ record.cost or 0.0 for record in records if record.status is not ValueStatus.EXACT
306
+ ),
307
+ latency_seconds=sum(record.latency_seconds or 0.0 for record in records),
308
+ status=status,
309
+ )