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.
tunarag/dataset.py ADDED
@@ -0,0 +1,481 @@
1
+ """Repeatable datasets and normalized evaluation-example ingestion."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import csv
6
+ import json
7
+ import math
8
+ import unicodedata
9
+ from collections.abc import AsyncIterator, Iterable, Mapping, Sequence
10
+ from dataclasses import dataclass, field, replace
11
+ from enum import Enum
12
+ from pathlib import Path
13
+ from types import MappingProxyType
14
+ from typing import Any
15
+
16
+ from .errors import DatasetError
17
+ from .serialization import content_hash
18
+
19
+ _SCHEMA_VERSION = 1
20
+
21
+
22
+ class DatasetSplit(str, Enum):
23
+ """Canonical purpose assigned to an evaluation example."""
24
+
25
+ OPTIMIZE = "optimize"
26
+ VALIDATION = "validation"
27
+ TEST = "test"
28
+
29
+
30
+ @dataclass(frozen=True, slots=True)
31
+ class SplitRatios:
32
+ """Deterministic split thresholds."""
33
+
34
+ optimize: float = 0.8
35
+ validation: float = 0.1
36
+ test: float = 0.1
37
+
38
+ def __post_init__(self) -> None:
39
+ values = (self.optimize, self.validation, self.test)
40
+ if any(
41
+ isinstance(value, bool) or not math.isfinite(value) or value < 0 for value in values
42
+ ):
43
+ raise ValueError("split ratios must be finite nonnegative numbers")
44
+ if not math.isclose(sum(values), 1.0, rel_tol=0.0, abs_tol=1e-12):
45
+ raise ValueError("split ratios must sum to 1")
46
+
47
+
48
+ @dataclass(frozen=True, slots=True)
49
+ class DatasetFieldMap:
50
+ """Map source record fields to the canonical evaluation schema."""
51
+
52
+ id: str = "id"
53
+ query: str = "query"
54
+ reference_answer: str = "reference_answer"
55
+ reference_contexts: str = "reference_contexts"
56
+ relevant_document_ids: str = "relevant_document_ids"
57
+ tags: str = "tags"
58
+ metadata: str = "metadata"
59
+ split: str = "split"
60
+
61
+ def __post_init__(self) -> None:
62
+ values = tuple(getattr(self, name) for name in self.__dataclass_fields__)
63
+ if any(not value.strip() for value in values):
64
+ raise ValueError("dataset field names must not be empty")
65
+ if len(set(values)) != len(values):
66
+ raise ValueError("dataset field names must be unique")
67
+
68
+
69
+ @dataclass(frozen=True, slots=True)
70
+ class EvaluationExample:
71
+ """One normalized, provenance-bearing evaluation example."""
72
+
73
+ id: str
74
+ query: str
75
+ reference_answer: str | None = None
76
+ reference_contexts: tuple[str, ...] = ()
77
+ relevant_document_ids: tuple[str, ...] = ()
78
+ tags: tuple[str, ...] = ()
79
+ metadata: Mapping[str, Any] = field(default_factory=dict)
80
+ split: DatasetSplit = DatasetSplit.OPTIMIZE
81
+ synthetic: bool = False
82
+
83
+ def __post_init__(self) -> None:
84
+ object.__setattr__(self, "id", _normalize_required(self.id, "example id"))
85
+ object.__setattr__(self, "query", _normalize_required(self.query, "query"))
86
+ if self.reference_answer is not None:
87
+ normalized_answer = _normalize_text(self.reference_answer)
88
+ object.__setattr__(self, "reference_answer", normalized_answer or None)
89
+ object.__setattr__(
90
+ self,
91
+ "reference_contexts",
92
+ _normalize_unique(self.reference_contexts, "reference context", unique=False),
93
+ )
94
+ object.__setattr__(
95
+ self,
96
+ "relevant_document_ids",
97
+ _normalize_unique(self.relevant_document_ids, "relevant document id"),
98
+ )
99
+ object.__setattr__(self, "tags", _normalize_unique(self.tags, "tag"))
100
+ normalized_metadata = _json_value(dict(self.metadata), path="metadata")
101
+ object.__setattr__(self, "metadata", MappingProxyType(normalized_metadata))
102
+ if not isinstance(self.split, DatasetSplit):
103
+ try:
104
+ object.__setattr__(self, "split", DatasetSplit(self.split))
105
+ except (TypeError, ValueError) as error:
106
+ raise ValueError("invalid dataset split") from error
107
+ if not isinstance(self.synthetic, bool):
108
+ raise TypeError("synthetic must be a boolean")
109
+
110
+ def semantic_data(self) -> dict[str, Any]:
111
+ """Return normalized content used for version hashing."""
112
+
113
+ return {
114
+ "id": self.id,
115
+ "query": self.query,
116
+ "reference_answer": self.reference_answer,
117
+ "reference_contexts": self.reference_contexts,
118
+ "relevant_document_ids": self.relevant_document_ids,
119
+ "tags": self.tags,
120
+ "metadata": self.metadata,
121
+ "split": self.split.value,
122
+ "synthetic": self.synthetic,
123
+ }
124
+
125
+
126
+ @dataclass(frozen=True, slots=True)
127
+ class EvaluationDataset:
128
+ """A normalized, immutable dataset with a content-derived version."""
129
+
130
+ examples: tuple[EvaluationExample, ...]
131
+ version: str = field(init=False)
132
+
133
+ def __post_init__(self) -> None:
134
+ if not self.examples:
135
+ raise ValueError("dataset must contain at least one example")
136
+ ids = [example.id for example in self.examples]
137
+ if len(ids) != len(set(ids)):
138
+ raise ValueError("dataset example ids must be unique")
139
+ version = content_hash(
140
+ {
141
+ "schema_version": _SCHEMA_VERSION,
142
+ "examples": [example.semantic_data() for example in self.examples],
143
+ },
144
+ namespace="tunarag:dataset:v1",
145
+ )
146
+ object.__setattr__(self, "version", version)
147
+
148
+ @classmethod
149
+ def from_records(
150
+ cls,
151
+ records: Iterable[Mapping[str, Any]],
152
+ *,
153
+ field_map: DatasetFieldMap | None = None,
154
+ source: str = "records",
155
+ ) -> EvaluationDataset:
156
+ """Normalize mapping records into an immutable dataset."""
157
+
158
+ mapping = field_map or DatasetFieldMap()
159
+ examples: list[EvaluationExample] = []
160
+ for index, record in enumerate(records, start=1):
161
+ if not isinstance(record, Mapping):
162
+ raise DatasetError(
163
+ "dataset record must be an object",
164
+ details={"source": source, "record": index},
165
+ )
166
+ try:
167
+ examples.append(_example_from_record(record, mapping))
168
+ except (TypeError, ValueError, json.JSONDecodeError) as error:
169
+ raise DatasetError(
170
+ "dataset record is invalid",
171
+ details={
172
+ "source": source,
173
+ "record": index,
174
+ "cause": type(error).__name__,
175
+ },
176
+ ) from error
177
+ try:
178
+ return cls(tuple(examples))
179
+ except ValueError as error:
180
+ raise DatasetError(
181
+ "dataset validation failed",
182
+ details={"source": source, "cause": type(error).__name__},
183
+ ) from error
184
+
185
+ @classmethod
186
+ def from_jsonl(
187
+ cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
188
+ ) -> EvaluationDataset:
189
+ """Load UTF-8 JSON Lines, ignoring blank lines."""
190
+
191
+ source = str(path)
192
+ records: list[Mapping[str, Any]] = []
193
+ try:
194
+ with Path(path).open("r", encoding="utf-8-sig") as stream:
195
+ for line_number, line in enumerate(stream, start=1):
196
+ if not line.strip():
197
+ continue
198
+ try:
199
+ value = json.loads(line)
200
+ except json.JSONDecodeError as error:
201
+ raise DatasetError(
202
+ "could not load JSONL dataset",
203
+ details={
204
+ "source": source,
205
+ "record": line_number,
206
+ "cause": type(error).__name__,
207
+ },
208
+ ) from error
209
+ if not isinstance(value, Mapping):
210
+ raise DatasetError(
211
+ "JSONL row must be an object",
212
+ details={"source": source, "record": line_number},
213
+ )
214
+ records.append(value)
215
+ except DatasetError:
216
+ raise
217
+ except (OSError, UnicodeError) as error:
218
+ raise DatasetError(
219
+ "could not load JSONL dataset",
220
+ details={"source": source, "cause": type(error).__name__},
221
+ ) from error
222
+ return cls.from_records(records, field_map=field_map, source=source)
223
+
224
+ @classmethod
225
+ def from_json(
226
+ cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
227
+ ) -> EvaluationDataset:
228
+ """Load a UTF-8 JSON array of objects."""
229
+
230
+ source = str(path)
231
+ try:
232
+ with Path(path).open("r", encoding="utf-8-sig") as stream:
233
+ value = json.load(stream)
234
+ except (OSError, UnicodeError, json.JSONDecodeError) as error:
235
+ raise DatasetError(
236
+ "could not load JSON dataset",
237
+ details={"source": source, "cause": type(error).__name__},
238
+ ) from error
239
+ if not isinstance(value, list):
240
+ raise DatasetError("JSON dataset must be an array", details={"source": source})
241
+ return cls.from_records(value, field_map=field_map, source=source)
242
+
243
+ @classmethod
244
+ def from_csv(
245
+ cls, path: str | Path, *, field_map: DatasetFieldMap | None = None
246
+ ) -> EvaluationDataset:
247
+ """Load a UTF-8 CSV with JSON-encoded collection fields."""
248
+
249
+ source = str(path)
250
+ try:
251
+ with Path(path).open("r", encoding="utf-8-sig", newline="") as stream:
252
+ records = list(csv.DictReader(stream))
253
+ except (OSError, UnicodeError, csv.Error) as error:
254
+ raise DatasetError(
255
+ "could not load CSV dataset",
256
+ details={"source": source, "cause": type(error).__name__},
257
+ ) from error
258
+ return cls.from_records(records, field_map=field_map, source=source)
259
+
260
+ def with_splits(
261
+ self,
262
+ *,
263
+ seed: int = 0,
264
+ ratios: SplitRatios | None = None,
265
+ group_metadata_key: str | None = None,
266
+ ) -> EvaluationDataset:
267
+ """Return a deterministically split copy, optionally keeping groups together."""
268
+
269
+ if isinstance(seed, bool) or not isinstance(seed, int):
270
+ raise TypeError("split seed must be an integer")
271
+ if group_metadata_key is not None and not group_metadata_key.strip():
272
+ raise ValueError("group metadata key must not be empty")
273
+ thresholds = ratios or SplitRatios()
274
+ examples = tuple(
275
+ replace(
276
+ example,
277
+ split=_select_split(
278
+ seed,
279
+ _group_identity(example, group_metadata_key),
280
+ thresholds,
281
+ ),
282
+ )
283
+ for example in self.examples
284
+ )
285
+ return EvaluationDataset(examples)
286
+
287
+ def select(self, split: DatasetSplit) -> EvaluationDataset:
288
+ """Return one nonempty canonical split."""
289
+
290
+ selected = tuple(example for example in self.examples if example.split is split)
291
+ if not selected:
292
+ raise DatasetError("selected dataset split is empty", details={"split": split.value})
293
+ return EvaluationDataset(selected)
294
+
295
+ def __aiter__(self) -> AsyncIterator[EvaluationExample]:
296
+ return self._iterate()
297
+
298
+ async def _iterate(self) -> AsyncIterator[EvaluationExample]:
299
+ for example in self.examples:
300
+ yield example
301
+
302
+
303
+ @dataclass(frozen=True, slots=True)
304
+ class InMemoryDataset:
305
+ """A versioned, repeatable dataset backed by an immutable tuple."""
306
+
307
+ examples: tuple[Any, ...]
308
+ version: str
309
+
310
+ def __post_init__(self) -> None:
311
+ if not self.examples:
312
+ raise ValueError("dataset must contain at least one example")
313
+ if not self.version.strip():
314
+ raise ValueError("dataset version must not be empty")
315
+
316
+ @classmethod
317
+ def from_iterable(cls, examples: Iterable[Any], *, version: str) -> InMemoryDataset:
318
+ """Materialize a repeatable dataset from an iterable."""
319
+
320
+ return cls(tuple(examples), version)
321
+
322
+ def __aiter__(self) -> AsyncIterator[Any]:
323
+ return self._iterate()
324
+
325
+ async def _iterate(self) -> AsyncIterator[Any]:
326
+ for example in self.examples:
327
+ yield example
328
+
329
+
330
+ def _example_from_record(record: Mapping[str, Any], fields: DatasetFieldMap) -> EvaluationExample:
331
+ query = _required_record_value(record, fields.query)
332
+ example_id = record.get(fields.id)
333
+ reference_answer = _optional_string(record.get(fields.reference_answer))
334
+ reference_contexts = _string_sequence(record.get(fields.reference_contexts))
335
+ relevant_document_ids = _string_sequence(record.get(fields.relevant_document_ids))
336
+ tags = _string_sequence(record.get(fields.tags))
337
+ metadata = _metadata(record.get(fields.metadata))
338
+ split = _split(record.get(fields.split))
339
+ if example_id is None or (isinstance(example_id, str) and not example_id.strip()):
340
+ example_id = content_hash(
341
+ {
342
+ "query": query,
343
+ "reference_answer": reference_answer,
344
+ "reference_contexts": reference_contexts,
345
+ "relevant_document_ids": relevant_document_ids,
346
+ "tags": tags,
347
+ "metadata": metadata,
348
+ "split": split.value,
349
+ "synthetic": False,
350
+ },
351
+ namespace="tunarag:example-id:v1",
352
+ )
353
+ if not isinstance(example_id, str):
354
+ raise TypeError("example id must be a string")
355
+ return EvaluationExample(
356
+ id=example_id,
357
+ query=query,
358
+ reference_answer=reference_answer,
359
+ reference_contexts=reference_contexts,
360
+ relevant_document_ids=relevant_document_ids,
361
+ tags=tags,
362
+ metadata=metadata,
363
+ split=split,
364
+ synthetic=False,
365
+ )
366
+
367
+
368
+ def _required_record_value(record: Mapping[str, Any], key: str) -> str:
369
+ value = record.get(key)
370
+ if not isinstance(value, str):
371
+ raise TypeError(f"{key} must be a string")
372
+ return value
373
+
374
+
375
+ def _optional_string(value: Any) -> str | None:
376
+ if value is None or value == "":
377
+ return None
378
+ if not isinstance(value, str):
379
+ raise TypeError("optional text field must be a string")
380
+ return value
381
+
382
+
383
+ def _string_sequence(value: Any) -> tuple[str, ...]:
384
+ if value is None or value == "":
385
+ return ()
386
+ if isinstance(value, str):
387
+ value = json.loads(value)
388
+ if not isinstance(value, Sequence) or isinstance(value, (str, bytes, bytearray)):
389
+ raise TypeError("collection field must be an array of strings")
390
+ if any(not isinstance(item, str) for item in value):
391
+ raise TypeError("collection field must contain only strings")
392
+ return tuple(value)
393
+
394
+
395
+ def _metadata(value: Any) -> Mapping[str, Any]:
396
+ if value is None or value == "":
397
+ return {}
398
+ if isinstance(value, str):
399
+ value = json.loads(value)
400
+ if not isinstance(value, Mapping):
401
+ raise TypeError("metadata must be an object")
402
+ return value
403
+
404
+
405
+ def _split(value: Any) -> DatasetSplit:
406
+ if value is None or value == "":
407
+ return DatasetSplit.OPTIMIZE
408
+ if not isinstance(value, str):
409
+ raise TypeError("split must be a string")
410
+ return DatasetSplit(value.strip().lower())
411
+
412
+
413
+ def _normalize_required(value: str, label: str) -> str:
414
+ normalized = _normalize_text(value)
415
+ if not normalized:
416
+ raise ValueError(f"{label} must not be empty")
417
+ return normalized
418
+
419
+
420
+ def _normalize_text(value: str) -> str:
421
+ if not isinstance(value, str):
422
+ raise TypeError("text value must be a string")
423
+ normalized = unicodedata.normalize("NFC", value.lstrip("\ufeff"))
424
+ normalized = normalized.replace("\r\n", "\n").replace("\r", "\n").strip()
425
+ if any(unicodedata.category(char) == "Cc" and char not in {"\n", "\t"} for char in normalized):
426
+ raise ValueError("text contains a disallowed control character")
427
+ return normalized
428
+
429
+
430
+ def _normalize_unique(values: Iterable[str], label: str, *, unique: bool = True) -> tuple[str, ...]:
431
+ normalized = tuple(_normalize_required(value, label) for value in values)
432
+ if unique and len(normalized) != len(set(normalized)):
433
+ raise ValueError(f"{label} values must be unique")
434
+ return normalized
435
+
436
+
437
+ def _json_value(value: Any, *, path: str) -> Any:
438
+ if value is None or isinstance(value, (bool, int, str)):
439
+ return value
440
+ if isinstance(value, float):
441
+ if not math.isfinite(value):
442
+ raise ValueError(f"{path} contains a nonfinite number")
443
+ return value
444
+ if isinstance(value, Mapping):
445
+ if any(not isinstance(key, str) for key in value):
446
+ raise TypeError(f"{path} keys must be strings")
447
+ return MappingProxyType(
448
+ {key: _json_value(item, path=f"{path}.{key}") for key, item in value.items()}
449
+ )
450
+ if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
451
+ return tuple(_json_value(item, path=f"{path}[]") for item in value)
452
+ raise TypeError(f"{path} contains unsupported value type {type(value).__name__}")
453
+
454
+
455
+ def _group_identity(example: EvaluationExample, metadata_key: str | None) -> str:
456
+ if metadata_key is None:
457
+ return example.id
458
+ if metadata_key not in example.metadata:
459
+ raise DatasetError(
460
+ "group metadata is missing",
461
+ details={"example_id": example.id, "metadata_key": metadata_key},
462
+ )
463
+ value = example.metadata[metadata_key]
464
+ if not isinstance(value, (str, int)) or isinstance(value, bool):
465
+ raise DatasetError(
466
+ "group metadata must be a string or integer",
467
+ details={"example_id": example.id, "metadata_key": metadata_key},
468
+ )
469
+ return str(value)
470
+
471
+
472
+ def _select_split(seed: int, identity: str, ratios: SplitRatios) -> DatasetSplit:
473
+ digest = content_hash(
474
+ {"seed": seed, "identity": identity}, namespace="tunarag:dataset-split:v1"
475
+ )
476
+ fraction = int(digest[:16], 16) / float(2**64)
477
+ if fraction < ratios.optimize:
478
+ return DatasetSplit.OPTIMIZE
479
+ if fraction < ratios.optimize + ratios.validation:
480
+ return DatasetSplit.VALIDATION
481
+ return DatasetSplit.TEST
tunarag/domain.py ADDED
@@ -0,0 +1,83 @@
1
+ """Immutable domain values shared by the package contracts."""
2
+
3
+ from collections.abc import Mapping
4
+ from dataclasses import dataclass
5
+ from enum import Enum
6
+ from typing import Any
7
+
8
+
9
+ class ValueStatus(str, Enum):
10
+ """Confidence state for an observed or derived value."""
11
+
12
+ EXACT = "exact"
13
+ ESTIMATED = "estimated"
14
+ UNAVAILABLE = "unavailable"
15
+
16
+
17
+ @dataclass(frozen=True, slots=True)
18
+ class Secret:
19
+ """A value whose contents are never exposed by representation or hashing."""
20
+
21
+ value: str
22
+
23
+ def __repr__(self) -> str:
24
+ return "Secret(***redacted***)"
25
+
26
+ def __str__(self) -> str:
27
+ return "***redacted***"
28
+
29
+
30
+ @dataclass(frozen=True, slots=True)
31
+ class MetricValue:
32
+ """A metric observation with its coverage and confidence state."""
33
+
34
+ name: str
35
+ value: float | None
36
+ status: ValueStatus = ValueStatus.EXACT
37
+ coverage: float = 1.0
38
+
39
+ def __post_init__(self) -> None:
40
+ if not self.name.strip():
41
+ raise ValueError("metric name must not be empty")
42
+ if not 0.0 <= self.coverage <= 1.0:
43
+ raise ValueError("metric coverage must be between 0 and 1")
44
+ if self.value is not None and not isinstance(self.value, (int, float)):
45
+ raise TypeError("metric value must be numeric or None")
46
+
47
+
48
+ @dataclass(frozen=True, slots=True)
49
+ class UsageRecord:
50
+ """Usage observed for one component of a trial."""
51
+
52
+ component: str
53
+ input_tokens: int | None = None
54
+ output_tokens: int | None = None
55
+ cost: float | None = None
56
+ latency_seconds: float | None = None
57
+ status: ValueStatus = ValueStatus.EXACT
58
+ pricing_version: str | None = None
59
+
60
+ def __post_init__(self) -> None:
61
+ if not self.component.strip():
62
+ raise ValueError("usage component must not be empty")
63
+ for name in ("input_tokens", "output_tokens"):
64
+ token_value: Any = getattr(self, name)
65
+ if token_value is not None and token_value < 0:
66
+ raise ValueError(f"{name} must not be negative")
67
+ for name in ("cost", "latency_seconds"):
68
+ usage_value: Any = getattr(self, name)
69
+ if usage_value is not None and usage_value < 0:
70
+ raise ValueError(f"{name} must not be negative")
71
+
72
+
73
+ @dataclass(frozen=True, slots=True)
74
+ class Candidate:
75
+ """A typed candidate configuration proposed for a trial."""
76
+
77
+ parameters: Mapping[str, Any]
78
+
79
+ def __post_init__(self) -> None:
80
+ if not self.parameters:
81
+ raise ValueError("candidate parameters must not be empty")
82
+ if any(not name.strip() for name in self.parameters):
83
+ raise ValueError("candidate parameter names must not be empty")