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/__init__.py ADDED
@@ -0,0 +1,166 @@
1
+ """Public package for typed RAG configuration optimization."""
2
+
3
+ from .cache import (
4
+ CacheEntry,
5
+ CacheKey,
6
+ CacheLookup,
7
+ CacheLookupStatus,
8
+ CachePrivacy,
9
+ SQLiteCache,
10
+ )
11
+ from .config import CachePolicy, Settings
12
+ from .dataset import (
13
+ DatasetFieldMap,
14
+ DatasetSplit,
15
+ EvaluationDataset,
16
+ EvaluationExample,
17
+ InMemoryDataset,
18
+ SplitRatios,
19
+ )
20
+ from .domain import Candidate, MetricValue, Secret, UsageRecord, ValueStatus
21
+ from .engine import Optimizer
22
+ from .errors import (
23
+ AdapterError,
24
+ BudgetError,
25
+ ConfigurationError,
26
+ DatasetError,
27
+ ErrorCode,
28
+ EvaluationError,
29
+ MissingOptionalDependencyError,
30
+ ResumeError,
31
+ SearchSpaceError,
32
+ StoreError,
33
+ SyncInAsyncContextError,
34
+ TunaRAGError,
35
+ )
36
+ from .evaluators import MetricEvaluator, MetricScorer, RagasEvaluator, RagasSampleMapper
37
+ from .integrations.runnables import LangChainAdapter, LangGraphAdapter
38
+ from .objective import Objective, ObjectiveDirection, ObjectiveResult, ObjectiveTerm
39
+ from .result import OptimizationResult, TrialSummary, UsageSummary
40
+ from .retry import RetryPolicy
41
+ from .search import (
42
+ Categorical,
43
+ FloatRange,
44
+ IntegerRange,
45
+ LLMSearch,
46
+ LLMSearchProvider,
47
+ LLMSuggestionRequest,
48
+ RandomSearch,
49
+ SearchObservation,
50
+ SearchSpace,
51
+ )
52
+ from .stopping import (
53
+ CostBudget,
54
+ MaxFailures,
55
+ MaxTrials,
56
+ NoImprovement,
57
+ StopContext,
58
+ TargetScore,
59
+ TimeBudget,
60
+ )
61
+ from .store import (
62
+ AttemptRecord,
63
+ AttemptStatus,
64
+ PersistedUsageRecord,
65
+ SQLiteStore,
66
+ StudyEvent,
67
+ StudyRecord,
68
+ StudyStatus,
69
+ TrialRecord,
70
+ TrialStatus,
71
+ )
72
+ from .synthetic import (
73
+ SourceDocument,
74
+ SourceSpan,
75
+ SyntheticDatasetBuilder,
76
+ SyntheticDatasetResult,
77
+ SyntheticGenerationRequest,
78
+ SyntheticGenerationResponse,
79
+ SyntheticQA,
80
+ SyntheticQAProvider,
81
+ )
82
+
83
+ __all__ = [
84
+ "AdapterError",
85
+ "AttemptRecord",
86
+ "AttemptStatus",
87
+ "BudgetError",
88
+ "CacheEntry",
89
+ "CacheKey",
90
+ "CacheLookup",
91
+ "CacheLookupStatus",
92
+ "CachePolicy",
93
+ "CachePrivacy",
94
+ "Categorical",
95
+ "Candidate",
96
+ "ConfigurationError",
97
+ "CostBudget",
98
+ "DatasetError",
99
+ "DatasetFieldMap",
100
+ "DatasetSplit",
101
+ "ErrorCode",
102
+ "EvaluationError",
103
+ "EvaluationDataset",
104
+ "EvaluationExample",
105
+ "FloatRange",
106
+ "InMemoryDataset",
107
+ "IntegerRange",
108
+ "LangChainAdapter",
109
+ "LangGraphAdapter",
110
+ "LLMSearch",
111
+ "LLMSearchProvider",
112
+ "LLMSuggestionRequest",
113
+ "MaxTrials",
114
+ "MaxFailures",
115
+ "MetricEvaluator",
116
+ "MetricScorer",
117
+ "MetricValue",
118
+ "MissingOptionalDependencyError",
119
+ "Objective",
120
+ "ObjectiveDirection",
121
+ "ObjectiveResult",
122
+ "ObjectiveTerm",
123
+ "NoImprovement",
124
+ "OptimizationResult",
125
+ "Optimizer",
126
+ "PersistedUsageRecord",
127
+ "ResumeError",
128
+ "RandomSearch",
129
+ "RagasEvaluator",
130
+ "RagasSampleMapper",
131
+ "RetryPolicy",
132
+ "Secret",
133
+ "SearchSpaceError",
134
+ "SearchObservation",
135
+ "SearchSpace",
136
+ "Settings",
137
+ "SQLiteCache",
138
+ "SQLiteStore",
139
+ "SourceDocument",
140
+ "SourceSpan",
141
+ "SplitRatios",
142
+ "StopContext",
143
+ "StoreError",
144
+ "StudyEvent",
145
+ "SyncInAsyncContextError",
146
+ "StudyRecord",
147
+ "StudyStatus",
148
+ "SyntheticDatasetBuilder",
149
+ "SyntheticDatasetResult",
150
+ "SyntheticGenerationRequest",
151
+ "SyntheticGenerationResponse",
152
+ "SyntheticQA",
153
+ "SyntheticQAProvider",
154
+ "TargetScore",
155
+ "TimeBudget",
156
+ "TrialRecord",
157
+ "TrialSummary",
158
+ "TrialStatus",
159
+ "TunaRAGError",
160
+ "UsageRecord",
161
+ "UsageSummary",
162
+ "ValueStatus",
163
+ "__version__",
164
+ ]
165
+
166
+ __version__ = "0.2.1"
tunarag/cache.py ADDED
@@ -0,0 +1,326 @@
1
+ """Versioned, checksum-verified local cache persistence."""
2
+
3
+ import hashlib
4
+ import math
5
+ import re
6
+ import sqlite3
7
+ import time
8
+ from collections.abc import Callable, Sequence
9
+ from dataclasses import dataclass
10
+ from enum import Enum
11
+ from pathlib import Path
12
+ from typing import Any, cast
13
+
14
+ from .serialization import content_hash
15
+
16
+ _NAMESPACE_PATTERN = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:-]{0,127}")
17
+ _DIGEST_PATTERN = re.compile(r"[0-9a-f]{64}")
18
+
19
+
20
+ class CachePrivacy(str, Enum):
21
+ """Privacy classification retained with a cache payload."""
22
+
23
+ STANDARD = "standard"
24
+ SENSITIVE = "sensitive"
25
+ RESTRICTED = "restricted"
26
+
27
+
28
+ class CacheLookupStatus(str, Enum):
29
+ """Outcome of a cache lookup."""
30
+
31
+ HIT = "hit"
32
+ MISS = "miss"
33
+ EXPIRED = "expired"
34
+ CORRUPT = "corrupt"
35
+ QUARANTINED = "quarantined"
36
+
37
+
38
+ class CacheEntryState(str, Enum):
39
+ """Persisted cache-entry lifecycle state."""
40
+
41
+ ACTIVE = "active"
42
+ EXPIRED = "expired"
43
+ QUARANTINED = "quarantined"
44
+
45
+
46
+ @dataclass(frozen=True, slots=True)
47
+ class CacheKey:
48
+ """Namespaced identity for one semantic cache value."""
49
+
50
+ namespace: str
51
+ digest: str
52
+ format_version: int = 1
53
+
54
+ def __post_init__(self) -> None:
55
+ if _NAMESPACE_PATTERN.fullmatch(self.namespace) is None:
56
+ raise ValueError("cache namespace must use 1-128 safe characters")
57
+ if _DIGEST_PATTERN.fullmatch(self.digest) is None:
58
+ raise ValueError("cache digest must be a lowercase SHA-256 hex digest")
59
+ if (
60
+ not isinstance(self.format_version, int)
61
+ or isinstance(self.format_version, bool)
62
+ or self.format_version < 1
63
+ ):
64
+ raise ValueError("cache format version must be a positive integer")
65
+
66
+ @classmethod
67
+ def from_parts(cls, namespace: str, *parts: Any, format_version: int = 1) -> "CacheKey":
68
+ """Build a key from canonical semantic inputs."""
69
+
70
+ if not parts:
71
+ raise ValueError("cache key requires at least one semantic part")
72
+ digest = content_hash(
73
+ {"format_version": format_version, "parts": parts},
74
+ namespace=f"tunarag:cache-key:{namespace}:v1",
75
+ )
76
+ return cls(namespace=namespace, digest=digest, format_version=format_version)
77
+
78
+
79
+ @dataclass(frozen=True, slots=True)
80
+ class CacheEntry:
81
+ """Metadata returned after an atomic cache write."""
82
+
83
+ key: CacheKey
84
+ checksum: str
85
+ privacy: CachePrivacy
86
+ created_at: float
87
+ expires_at: float | None
88
+
89
+
90
+ @dataclass(frozen=True, slots=True)
91
+ class CacheLookup:
92
+ """Payload and metadata returned by a cache lookup."""
93
+
94
+ key: CacheKey
95
+ status: CacheLookupStatus
96
+ payload: bytes | None = None
97
+ checksum: str | None = None
98
+ privacy: CachePrivacy | None = None
99
+ created_at: float | None = None
100
+ expires_at: float | None = None
101
+ reason: str | None = None
102
+
103
+ @property
104
+ def is_hit(self) -> bool:
105
+ """Return whether this lookup contains a verified payload."""
106
+
107
+ return self.status is CacheLookupStatus.HIT
108
+
109
+
110
+ class SQLiteCache:
111
+ """Transactional SQLite cache with expiry and corruption quarantine."""
112
+
113
+ def __init__(
114
+ self,
115
+ path: str | Path = ":memory:",
116
+ *,
117
+ clock: Callable[[], float] = time.time,
118
+ ) -> None:
119
+ self._connection = sqlite3.connect(str(path))
120
+ self._connection.row_factory = sqlite3.Row
121
+ self._clock = clock
122
+ self.initialize()
123
+
124
+ def initialize(self) -> None:
125
+ """Create the independently versioned cache schema."""
126
+
127
+ self._connection.executescript(
128
+ """
129
+ CREATE TABLE IF NOT EXISTS cache_schema_version (
130
+ version INTEGER PRIMARY KEY
131
+ );
132
+ INSERT OR IGNORE INTO cache_schema_version(version) VALUES (1);
133
+ CREATE TABLE IF NOT EXISTS cache_entries (
134
+ namespace TEXT NOT NULL,
135
+ format_version INTEGER NOT NULL,
136
+ key_digest TEXT NOT NULL,
137
+ payload BLOB NOT NULL,
138
+ checksum TEXT NOT NULL,
139
+ privacy TEXT NOT NULL,
140
+ state TEXT NOT NULL,
141
+ created_at REAL NOT NULL,
142
+ expires_at REAL,
143
+ last_accessed_at REAL,
144
+ hit_count INTEGER NOT NULL DEFAULT 0,
145
+ quarantine_reason TEXT,
146
+ PRIMARY KEY(namespace, format_version, key_digest),
147
+ CHECK(format_version > 0),
148
+ CHECK(hit_count >= 0),
149
+ CHECK(privacy IN ('standard', 'sensitive', 'restricted')),
150
+ CHECK(state IN ('active', 'expired', 'quarantined'))
151
+ );
152
+ CREATE INDEX IF NOT EXISTS idx_cache_entries_retention
153
+ ON cache_entries(state, expires_at);
154
+ """
155
+ )
156
+ self._connection.commit()
157
+ versions = [
158
+ int(row[0])
159
+ for row in self._connection.execute(
160
+ "SELECT version FROM cache_schema_version ORDER BY version"
161
+ )
162
+ ]
163
+ if versions != [1]:
164
+ self._connection.close()
165
+ raise RuntimeError(f"unsupported cache schema versions: {versions}")
166
+
167
+ def put(
168
+ self,
169
+ key: CacheKey,
170
+ payload: bytes | bytearray | memoryview,
171
+ *,
172
+ privacy: CachePrivacy = CachePrivacy.STANDARD,
173
+ ttl_seconds: float | None = None,
174
+ ) -> CacheEntry:
175
+ """Atomically insert or replace a cache payload."""
176
+
177
+ payload_bytes = bytes(payload)
178
+ if ttl_seconds is not None:
179
+ if not math.isfinite(ttl_seconds) or ttl_seconds <= 0:
180
+ raise ValueError("cache TTL must be a positive finite number")
181
+ now = self._now()
182
+ expires_at = None if ttl_seconds is None else now + ttl_seconds
183
+ checksum = _payload_checksum(payload_bytes)
184
+ with self._connection:
185
+ self._connection.execute(
186
+ """
187
+ INSERT INTO cache_entries(
188
+ namespace, format_version, key_digest, payload, checksum, privacy,
189
+ state, created_at, expires_at, last_accessed_at, hit_count,
190
+ quarantine_reason
191
+ ) VALUES (?, ?, ?, ?, ?, ?, 'active', ?, ?, NULL, 0, NULL)
192
+ ON CONFLICT(namespace, format_version, key_digest) DO UPDATE SET
193
+ payload = excluded.payload,
194
+ checksum = excluded.checksum,
195
+ privacy = excluded.privacy,
196
+ state = 'active',
197
+ created_at = excluded.created_at,
198
+ expires_at = excluded.expires_at,
199
+ last_accessed_at = NULL,
200
+ hit_count = 0,
201
+ quarantine_reason = NULL
202
+ """,
203
+ (
204
+ key.namespace,
205
+ key.format_version,
206
+ key.digest,
207
+ payload_bytes,
208
+ checksum,
209
+ privacy.value,
210
+ now,
211
+ expires_at,
212
+ ),
213
+ )
214
+ return CacheEntry(key, checksum, privacy, now, expires_at)
215
+
216
+ def get(self, key: CacheKey) -> CacheLookup:
217
+ """Return a verified payload or an explicit miss reason."""
218
+
219
+ row = self._select(key)
220
+ if row is None:
221
+ return CacheLookup(key, CacheLookupStatus.MISS)
222
+
223
+ privacy = CachePrivacy(row["privacy"])
224
+ created_at = float(row["created_at"])
225
+ expires_at = row["expires_at"]
226
+ metadata = {
227
+ "checksum": row["checksum"],
228
+ "privacy": privacy,
229
+ "created_at": created_at,
230
+ "expires_at": expires_at,
231
+ }
232
+ state = CacheEntryState(row["state"])
233
+ if state is CacheEntryState.QUARANTINED:
234
+ return CacheLookup(
235
+ key,
236
+ CacheLookupStatus.QUARANTINED,
237
+ reason=row["quarantine_reason"],
238
+ **metadata,
239
+ )
240
+
241
+ now = self._now()
242
+ if state is CacheEntryState.EXPIRED or (expires_at is not None and expires_at <= now):
243
+ if state is CacheEntryState.ACTIVE:
244
+ with self._connection:
245
+ self._connection.execute(
246
+ """UPDATE cache_entries SET state = 'expired'
247
+ WHERE namespace = ? AND format_version = ? AND key_digest = ?""",
248
+ (key.namespace, key.format_version, key.digest),
249
+ )
250
+ return CacheLookup(key, CacheLookupStatus.EXPIRED, **metadata)
251
+
252
+ payload = bytes(row["payload"])
253
+ if _payload_checksum(payload) != row["checksum"]:
254
+ reason = "payload checksum mismatch"
255
+ with self._connection:
256
+ self._connection.execute(
257
+ """UPDATE cache_entries
258
+ SET state = 'quarantined', quarantine_reason = ?
259
+ WHERE namespace = ? AND format_version = ? AND key_digest = ?""",
260
+ (reason, key.namespace, key.format_version, key.digest),
261
+ )
262
+ return CacheLookup(key, CacheLookupStatus.CORRUPT, reason=reason, **metadata)
263
+
264
+ with self._connection:
265
+ self._connection.execute(
266
+ """UPDATE cache_entries
267
+ SET last_accessed_at = ?, hit_count = hit_count + 1
268
+ WHERE namespace = ? AND format_version = ? AND key_digest = ?""",
269
+ (now, key.namespace, key.format_version, key.digest),
270
+ )
271
+ return CacheLookup(key, CacheLookupStatus.HIT, payload=payload, **metadata)
272
+
273
+ def delete(self, key: CacheKey) -> bool:
274
+ """Delete one cache entry and report whether it existed."""
275
+
276
+ with self._connection:
277
+ cursor = self._connection.execute(
278
+ """DELETE FROM cache_entries
279
+ WHERE namespace = ? AND format_version = ? AND key_digest = ?""",
280
+ (key.namespace, key.format_version, key.digest),
281
+ )
282
+ return cursor.rowcount == 1
283
+
284
+ def prune(self, *, include_quarantined: bool = False) -> int:
285
+ """Remove expired entries and optionally quarantined entries."""
286
+
287
+ now = self._now()
288
+ conditions: Sequence[str] = (
289
+ "state = 'expired'",
290
+ "(expires_at IS NOT NULL AND expires_at <= ?)",
291
+ )
292
+ where = " OR ".join(conditions)
293
+ parameters: list[float] = [now]
294
+ if include_quarantined:
295
+ where += " OR state = 'quarantined'"
296
+ with self._connection:
297
+ cursor = self._connection.execute(
298
+ f"DELETE FROM cache_entries WHERE {where}", parameters
299
+ )
300
+ return cursor.rowcount
301
+
302
+ def close(self) -> None:
303
+ """Close the underlying SQLite connection."""
304
+
305
+ self._connection.close()
306
+
307
+ def _select(self, key: CacheKey) -> sqlite3.Row | None:
308
+ return cast(
309
+ sqlite3.Row | None,
310
+ self._connection.execute(
311
+ """SELECT payload, checksum, privacy, state, created_at, expires_at,
312
+ quarantine_reason FROM cache_entries
313
+ WHERE namespace = ? AND format_version = ? AND key_digest = ?""",
314
+ (key.namespace, key.format_version, key.digest),
315
+ ).fetchone(),
316
+ )
317
+
318
+ def _now(self) -> float:
319
+ now = float(self._clock())
320
+ if not math.isfinite(now):
321
+ raise ValueError("cache clock must return a finite timestamp")
322
+ return now
323
+
324
+
325
+ def _payload_checksum(payload: bytes) -> str:
326
+ return hashlib.sha256(payload).hexdigest()