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/synthetic.py ADDED
@@ -0,0 +1,257 @@
1
+ """Provider-neutral synthetic QA dataset generation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import math
7
+ import unicodedata
8
+ from collections.abc import Mapping, Sequence
9
+ from dataclasses import dataclass, field
10
+ from typing import Any, Protocol
11
+
12
+ from .contracts import Fingerprintable
13
+ from .dataset import EvaluationDataset, EvaluationExample
14
+ from .domain import UsageRecord
15
+ from .errors import DatasetError
16
+ from .serialization import content_hash
17
+
18
+
19
+ @dataclass(frozen=True, slots=True)
20
+ class SourceDocument:
21
+ """One source document approved for synthetic generation."""
22
+
23
+ id: str
24
+ text: str
25
+ metadata: Mapping[str, Any] = field(default_factory=dict)
26
+
27
+ def __post_init__(self) -> None:
28
+ object.__setattr__(self, "id", _required_text(self.id, "source document id"))
29
+ object.__setattr__(self, "text", _required_text(self.text, "source document text"))
30
+
31
+
32
+ @dataclass(frozen=True, slots=True)
33
+ class SourceSpan:
34
+ """Half-open source character range supporting a generated QA pair."""
35
+
36
+ start: int
37
+ end: int
38
+
39
+ def __post_init__(self) -> None:
40
+ if any(
41
+ isinstance(value, bool) or not isinstance(value, int)
42
+ for value in (self.start, self.end)
43
+ ):
44
+ raise TypeError("source span offsets must be integers")
45
+ if self.start < 0 or self.end <= self.start:
46
+ raise ValueError("source span must be a nonempty half-open range")
47
+
48
+
49
+ @dataclass(frozen=True, slots=True)
50
+ class SyntheticQA:
51
+ """One provider-produced question, answer, and supporting provenance."""
52
+
53
+ query: str
54
+ answer: str
55
+ source_spans: tuple[SourceSpan, ...] = ()
56
+ confidence: float | None = None
57
+
58
+ def __post_init__(self) -> None:
59
+ object.__setattr__(self, "query", _required_text(self.query, "synthetic query"))
60
+ object.__setattr__(self, "answer", _required_text(self.answer, "synthetic answer"))
61
+ if self.confidence is not None:
62
+ if (
63
+ isinstance(self.confidence, bool)
64
+ or not math.isfinite(self.confidence)
65
+ or not 0.0 <= self.confidence <= 1.0
66
+ ):
67
+ raise ValueError("synthetic confidence must be between 0 and 1")
68
+
69
+
70
+ @dataclass(frozen=True, slots=True)
71
+ class SyntheticGenerationRequest:
72
+ """Structured request sent to a synthetic QA provider."""
73
+
74
+ document: SourceDocument
75
+ count: int
76
+ seed: int
77
+
78
+
79
+ @dataclass(frozen=True, slots=True)
80
+ class SyntheticGenerationResponse:
81
+ """Provider output with exact or estimated usage records."""
82
+
83
+ items: tuple[SyntheticQA, ...]
84
+ usage: tuple[UsageRecord, ...] = ()
85
+
86
+
87
+ class SyntheticQAProvider(Protocol):
88
+ """Generate typed QA pairs for one source document."""
89
+
90
+ async def generate(
91
+ self, request: SyntheticGenerationRequest
92
+ ) -> SyntheticGenerationResponse: ...
93
+
94
+
95
+ @dataclass(frozen=True, slots=True)
96
+ class SyntheticDatasetResult:
97
+ """Generated evaluation dataset and provider usage."""
98
+
99
+ dataset: EvaluationDataset
100
+ usage: tuple[UsageRecord, ...]
101
+
102
+
103
+ class SyntheticDatasetBuilder:
104
+ """Generate a validated dataset with bounded document concurrency."""
105
+
106
+ def __init__(
107
+ self,
108
+ provider: SyntheticQAProvider,
109
+ *,
110
+ samples_per_document: int = 3,
111
+ concurrency: int = 4,
112
+ seed: int = 0,
113
+ generator_version: str = "unspecified",
114
+ ) -> None:
115
+ for name, value in (
116
+ ("samples_per_document", samples_per_document),
117
+ ("concurrency", concurrency),
118
+ ):
119
+ if isinstance(value, bool) or not isinstance(value, int) or value < 1:
120
+ raise ValueError(f"{name} must be a positive integer")
121
+ if isinstance(seed, bool) or not isinstance(seed, int):
122
+ raise TypeError("synthetic generation seed must be an integer")
123
+ self._provider = provider
124
+ self._samples_per_document = samples_per_document
125
+ self._concurrency = concurrency
126
+ self._seed = seed
127
+ self._generator_version = _required_text(generator_version, "generator version")
128
+
129
+ async def generate(self, documents: Sequence[SourceDocument]) -> SyntheticDatasetResult:
130
+ """Generate and validate examples in source-document order."""
131
+
132
+ if not documents:
133
+ raise DatasetError("synthetic generation requires source documents")
134
+ if len({document.id for document in documents}) != len(documents):
135
+ raise DatasetError("source document ids must be unique")
136
+ semaphore = asyncio.Semaphore(self._concurrency)
137
+
138
+ async def bounded(document: SourceDocument) -> SyntheticGenerationResponse:
139
+ async with semaphore:
140
+ request = SyntheticGenerationRequest(
141
+ document,
142
+ self._samples_per_document,
143
+ _document_seed(self._seed, document.id),
144
+ )
145
+ try:
146
+ response = await self._provider.generate(request)
147
+ except Exception as error:
148
+ raise DatasetError(
149
+ "synthetic QA provider failed",
150
+ details={"document_id": document.id, "cause": type(error).__name__},
151
+ ) from error
152
+ return _validate_response(response, request)
153
+
154
+ responses = await asyncio.gather(*(bounded(document) for document in documents))
155
+ examples: list[EvaluationExample] = []
156
+ usage: list[UsageRecord] = []
157
+ provider_identity = _provider_identity(self._provider)
158
+ for document, response in zip(documents, responses, strict=True):
159
+ request_seed = _document_seed(self._seed, document.id)
160
+ usage.extend(response.usage)
161
+ for index, item in enumerate(response.items):
162
+ contexts = tuple(document.text[span.start : span.end] for span in item.source_spans)
163
+ if not contexts:
164
+ contexts = (document.text,)
165
+ metadata = {
166
+ "source_document_id": document.id,
167
+ "source_document_metadata": document.metadata,
168
+ "source_spans": [
169
+ {"start": span.start, "end": span.end} for span in item.source_spans
170
+ ],
171
+ "generator": provider_identity,
172
+ "generator_version": self._generator_version,
173
+ "seed": request_seed,
174
+ "confidence": item.confidence,
175
+ "review_status": "unreviewed",
176
+ }
177
+ example_id = content_hash(
178
+ {
179
+ "document_id": document.id,
180
+ "index": index,
181
+ "qa": metadata,
182
+ "query": item.query,
183
+ "answer": item.answer,
184
+ },
185
+ namespace="tunarag:synthetic-example:v1",
186
+ )
187
+ examples.append(
188
+ EvaluationExample(
189
+ id=example_id,
190
+ query=item.query,
191
+ reference_answer=item.answer,
192
+ reference_contexts=contexts,
193
+ relevant_document_ids=(document.id,),
194
+ tags=("synthetic",),
195
+ metadata=metadata,
196
+ synthetic=True,
197
+ )
198
+ )
199
+ return SyntheticDatasetResult(EvaluationDataset(tuple(examples)), tuple(usage))
200
+
201
+
202
+ def _validate_response(
203
+ response: Any, request: SyntheticGenerationRequest
204
+ ) -> SyntheticGenerationResponse:
205
+ if not isinstance(response, SyntheticGenerationResponse):
206
+ raise DatasetError(
207
+ "synthetic provider returned an invalid response",
208
+ details={"document_id": request.document.id},
209
+ )
210
+ if len(response.items) != request.count:
211
+ raise DatasetError(
212
+ "synthetic provider returned an unexpected item count",
213
+ details={
214
+ "document_id": request.document.id,
215
+ "expected": request.count,
216
+ "actual": len(response.items),
217
+ },
218
+ )
219
+ if any(not isinstance(item, SyntheticQA) for item in response.items):
220
+ raise DatasetError("synthetic provider returned invalid QA items")
221
+ if any(not isinstance(item, UsageRecord) for item in response.usage):
222
+ raise DatasetError("synthetic provider returned invalid usage records")
223
+ for item in response.items:
224
+ if any(span.end > len(request.document.text) for span in item.source_spans):
225
+ raise DatasetError(
226
+ "synthetic source span exceeds its document",
227
+ details={"document_id": request.document.id},
228
+ )
229
+ return response
230
+
231
+
232
+ def _document_seed(seed: int, document_id: str) -> int:
233
+ digest = content_hash(
234
+ {"seed": seed, "document_id": document_id}, namespace="tunarag:synthetic-seed:v1"
235
+ )
236
+ return int(digest[:16], 16)
237
+
238
+
239
+ def _provider_identity(provider: SyntheticQAProvider) -> Mapping[str, Any]:
240
+ identity: dict[str, Any] = {
241
+ "type": f"{type(provider).__module__}.{type(provider).__qualname__}"
242
+ }
243
+ if isinstance(provider, Fingerprintable):
244
+ identity["fingerprint"] = provider.fingerprint()
245
+ return identity
246
+
247
+
248
+ def _required_text(value: str, label: str) -> str:
249
+ if not isinstance(value, str):
250
+ raise TypeError(f"{label} must be a string")
251
+ normalized = unicodedata.normalize("NFC", value.lstrip("\ufeff"))
252
+ normalized = normalized.replace("\r\n", "\n").replace("\r", "\n").strip()
253
+ if not normalized:
254
+ raise ValueError(f"{label} must not be empty")
255
+ if any(unicodedata.category(char) == "Cc" and char not in {"\n", "\t"} for char in normalized):
256
+ raise ValueError(f"{label} contains a disallowed control character")
257
+ return normalized