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/retry.py ADDED
@@ -0,0 +1,70 @@
1
+ """Deterministic retry and timeout policy."""
2
+
3
+ import asyncio
4
+ import math
5
+ from dataclasses import dataclass
6
+
7
+ from .errors import TunaRAGError
8
+
9
+
10
+ @dataclass(frozen=True, slots=True)
11
+ class RetryPolicy:
12
+ """Bound retries, exponential backoff, and execution timeouts."""
13
+
14
+ max_retries: int = 2
15
+ sample_timeout_seconds: float | None = 60.0
16
+ trial_timeout_seconds: float | None = 300.0
17
+ base_delay_seconds: float = 0.5
18
+ max_delay_seconds: float = 30.0
19
+
20
+ def __post_init__(self) -> None:
21
+ if (
22
+ not isinstance(self.max_retries, int)
23
+ or isinstance(self.max_retries, bool)
24
+ or self.max_retries < 0
25
+ ):
26
+ raise ValueError("max_retries must be a nonnegative integer")
27
+ _validate_optional_timeout("sample timeout", self.sample_timeout_seconds)
28
+ _validate_optional_timeout("trial timeout", self.trial_timeout_seconds)
29
+ _validate_nonnegative_finite("base retry delay", self.base_delay_seconds)
30
+ _validate_nonnegative_finite("maximum retry delay", self.max_delay_seconds)
31
+ if self.base_delay_seconds > self.max_delay_seconds:
32
+ raise ValueError("base retry delay must not exceed maximum retry delay")
33
+
34
+ @property
35
+ def max_attempts(self) -> int:
36
+ """Return the initial attempt plus configured retries."""
37
+
38
+ return self.max_retries + 1
39
+
40
+ @property
41
+ def lease_seconds(self) -> float:
42
+ """Return a recovery lease longer than the configured trial timeout."""
43
+
44
+ return 86_400.0 if self.trial_timeout_seconds is None else self.trial_timeout_seconds + 60.0
45
+
46
+ def delay_before_retry(self, failed_attempt_number: int) -> float:
47
+ """Return deterministic capped exponential backoff."""
48
+
49
+ if failed_attempt_number < 0:
50
+ raise ValueError("failed attempt number must not be negative")
51
+ return float(
52
+ min(self.base_delay_seconds * (2**failed_attempt_number), self.max_delay_seconds)
53
+ )
54
+
55
+ def is_retryable(self, error: BaseException) -> bool:
56
+ """Classify known transient failures without parsing messages."""
57
+
58
+ if isinstance(error, TunaRAGError):
59
+ return error.retryable
60
+ return isinstance(error, (TimeoutError, asyncio.TimeoutError, ConnectionError))
61
+
62
+
63
+ def _validate_optional_timeout(name: str, value: float | None) -> None:
64
+ if value is not None and (not math.isfinite(value) or value <= 0):
65
+ raise ValueError(f"{name} must be a positive finite number or None")
66
+
67
+
68
+ def _validate_nonnegative_finite(name: str, value: float) -> None:
69
+ if not math.isfinite(value) or value < 0:
70
+ raise ValueError(f"{name} must be a nonnegative finite number")
tunarag/search.py ADDED
@@ -0,0 +1,286 @@
1
+ """Typed search spaces and the deterministic random strategy."""
2
+
3
+ import math
4
+ import random
5
+ from collections.abc import Mapping, Sequence
6
+ from dataclasses import dataclass
7
+ from typing import Any, Protocol, cast
8
+
9
+ from .contracts import Fingerprintable
10
+ from .domain import Candidate, MetricValue, Secret, UsageRecord
11
+ from .errors import SearchSpaceError, redact
12
+
13
+
14
+ class Parameter(Protocol):
15
+ """Sample one typed parameter value from a random stream."""
16
+
17
+ def sample(self, generator: random.Random) -> Any: ...
18
+
19
+ def accepts(self, value: Any) -> bool: ...
20
+
21
+ def describe(self) -> Mapping[str, Any]: ...
22
+
23
+
24
+ @dataclass(frozen=True, slots=True)
25
+ class IntegerRange:
26
+ """Inclusive integer range."""
27
+
28
+ minimum: int
29
+ maximum: int
30
+
31
+ def __post_init__(self) -> None:
32
+ if any(
33
+ isinstance(value, bool) or not isinstance(value, int)
34
+ for value in (self.minimum, self.maximum)
35
+ ):
36
+ raise TypeError("integer range bounds must be integers")
37
+ if self.minimum > self.maximum:
38
+ raise ValueError("integer range minimum must not exceed maximum")
39
+
40
+ def sample(self, generator: random.Random) -> int:
41
+ return generator.randint(self.minimum, self.maximum)
42
+
43
+ def accepts(self, value: Any) -> bool:
44
+ return (
45
+ isinstance(value, int)
46
+ and not isinstance(value, bool)
47
+ and self.minimum <= value <= self.maximum
48
+ )
49
+
50
+ def describe(self) -> Mapping[str, Any]:
51
+ return {"type": "integer", "minimum": self.minimum, "maximum": self.maximum}
52
+
53
+
54
+ @dataclass(frozen=True, slots=True)
55
+ class FloatRange:
56
+ """Continuous float range."""
57
+
58
+ minimum: float
59
+ maximum: float
60
+
61
+ def __post_init__(self) -> None:
62
+ if any(
63
+ isinstance(value, bool)
64
+ or not isinstance(value, (int, float))
65
+ or not math.isfinite(value)
66
+ for value in (self.minimum, self.maximum)
67
+ ):
68
+ raise ValueError("float range bounds must be finite numbers")
69
+ if self.minimum > self.maximum:
70
+ raise ValueError("float range minimum must not exceed maximum")
71
+
72
+ def sample(self, generator: random.Random) -> float:
73
+ return generator.uniform(self.minimum, self.maximum)
74
+
75
+ def accepts(self, value: Any) -> bool:
76
+ return (
77
+ isinstance(value, (int, float))
78
+ and not isinstance(value, bool)
79
+ and self.minimum <= value <= self.maximum
80
+ )
81
+
82
+ def describe(self) -> Mapping[str, Any]:
83
+ return {"type": "number", "minimum": self.minimum, "maximum": self.maximum}
84
+
85
+
86
+ @dataclass(frozen=True, slots=True)
87
+ class Categorical:
88
+ """A non-empty collection of candidate values."""
89
+
90
+ values: tuple[Any, ...]
91
+
92
+ def __post_init__(self) -> None:
93
+ if not self.values:
94
+ raise ValueError("categorical values must not be empty")
95
+
96
+ def sample(self, generator: random.Random) -> Any:
97
+ return generator.choice(self.values)
98
+
99
+ def accepts(self, value: Any) -> bool:
100
+ return value in self.values
101
+
102
+ def describe(self) -> Mapping[str, Any]:
103
+ return {
104
+ "type": "categorical",
105
+ "values": tuple(value for value in self.values if not isinstance(value, Secret)),
106
+ }
107
+
108
+
109
+ @dataclass(frozen=True, slots=True)
110
+ class SearchSpace:
111
+ """Named typed parameters from which candidates can be sampled."""
112
+
113
+ parameters: dict[str, Parameter]
114
+
115
+ def __post_init__(self) -> None:
116
+ if not self.parameters:
117
+ raise ValueError("search space must contain at least one parameter")
118
+ if any(not name.strip() for name in self.parameters):
119
+ raise ValueError("search-space parameter names must not be empty")
120
+ object.__setattr__(self, "parameters", dict(self.parameters))
121
+
122
+ def sample(self, generator: random.Random) -> Candidate:
123
+ return Candidate(
124
+ {name: parameter.sample(generator) for name, parameter in self.parameters.items()}
125
+ )
126
+
127
+ def validate(self, parameters: Mapping[str, Any]) -> Candidate:
128
+ """Validate an exact candidate mapping against every parameter."""
129
+
130
+ if set(parameters) != set(self.parameters):
131
+ raise ValueError("candidate keys must exactly match the search space")
132
+ for name, parameter in self.parameters.items():
133
+ if not parameter.accepts(parameters[name]):
134
+ raise ValueError(f"candidate value is outside the search space: {name}")
135
+ return Candidate(dict(parameters))
136
+
137
+ def describe(self) -> Mapping[str, Mapping[str, Any]]:
138
+ """Return a provider-safe structured search-space description."""
139
+
140
+ return {name: parameter.describe() for name, parameter in self.parameters.items()}
141
+
142
+
143
+ class RandomSearch:
144
+ """Seeded random strategy implementing the common search contract."""
145
+
146
+ def __init__(self, space: SearchSpace, *, seed: int | None = None) -> None:
147
+ self._space = space
148
+ self._seed = seed
149
+ self._generator = random.Random(seed)
150
+
151
+ async def suggest(self) -> Candidate:
152
+ return self._space.sample(self._generator)
153
+
154
+ async def observe(self, candidate: Candidate, metrics: Any, usage: Any) -> None:
155
+ del candidate, metrics, usage
156
+
157
+ async def replay(self, candidate: Candidate, metrics: Any, usage: Any) -> None:
158
+ """Advance and verify deterministic generator state during resume."""
159
+
160
+ del metrics, usage
161
+ expected = self._space.sample(self._generator)
162
+ if expected != candidate:
163
+ raise ValueError("persisted candidate does not match seeded random search state")
164
+
165
+ def fingerprint(self) -> Any:
166
+ """Return stable search-space and seed semantics."""
167
+
168
+ return {"space": self._space, "seed": self._seed}
169
+
170
+
171
+ @dataclass(frozen=True, slots=True)
172
+ class SearchObservation:
173
+ """A prior reserved candidate and any committed metrics."""
174
+
175
+ candidate: Candidate
176
+ metrics: tuple[MetricValue, ...]
177
+
178
+
179
+ @dataclass(frozen=True, slots=True)
180
+ class LLMSuggestionRequest:
181
+ """Provider-neutral structured input for one suggestion attempt."""
182
+
183
+ search_space: Mapping[str, Mapping[str, Any]]
184
+ observations: tuple[SearchObservation, ...]
185
+ suggestion_number: int
186
+ attempt_number: int
187
+
188
+
189
+ class LLMSearchProvider(Protocol):
190
+ """Produce one structured candidate proposal without provider coupling."""
191
+
192
+ async def suggest(self, request: LLMSuggestionRequest) -> Mapping[str, Any]: ...
193
+
194
+
195
+ class LLMSearch:
196
+ """Validate bounded provider suggestions with deterministic random fallback."""
197
+
198
+ def __init__(
199
+ self,
200
+ space: SearchSpace,
201
+ provider: LLMSearchProvider,
202
+ *,
203
+ max_suggestion_attempts: int = 3,
204
+ fallback_seed: int = 0,
205
+ fallback: bool = True,
206
+ ) -> None:
207
+ if (
208
+ not isinstance(max_suggestion_attempts, int)
209
+ or isinstance(max_suggestion_attempts, bool)
210
+ or max_suggestion_attempts < 1
211
+ ):
212
+ raise ValueError("maximum suggestion attempts must be positive")
213
+ self._space = space
214
+ self._provider = provider
215
+ self._max_suggestion_attempts = max_suggestion_attempts
216
+ self._fallback_seed = fallback_seed
217
+ self._fallback = fallback
218
+ self._suggestion_number = 0
219
+ self._observations: list[SearchObservation] = []
220
+
221
+ async def suggest(self) -> Candidate:
222
+ causes: list[str] = []
223
+ for attempt_number in range(self._max_suggestion_attempts):
224
+ request = LLMSuggestionRequest(
225
+ self._space.describe(),
226
+ tuple(_provider_observation(item) for item in self._observations),
227
+ self._suggestion_number,
228
+ attempt_number,
229
+ )
230
+ try:
231
+ proposal = await self._provider.suggest(request)
232
+ candidate = self._space.validate(proposal)
233
+ except Exception as error:
234
+ causes.append(type(error).__name__)
235
+ continue
236
+ self._suggestion_number += 1
237
+ return candidate
238
+ if self._fallback:
239
+ generator = random.Random(f"{self._fallback_seed}:{self._suggestion_number}")
240
+ candidate = self._space.sample(generator)
241
+ self._suggestion_number += 1
242
+ return candidate
243
+ raise SearchSpaceError(
244
+ "LLM search exhausted invalid suggestion attempts",
245
+ stage="strategy.suggest",
246
+ details={"causes": causes, "attempts": self._max_suggestion_attempts},
247
+ )
248
+
249
+ async def observe(
250
+ self,
251
+ candidate: Candidate,
252
+ metrics: Sequence[MetricValue],
253
+ usage: Sequence[UsageRecord],
254
+ ) -> None:
255
+ del usage
256
+ self._observations.append(SearchObservation(candidate, tuple(metrics)))
257
+
258
+ async def replay(
259
+ self,
260
+ candidate: Candidate,
261
+ metrics: Sequence[MetricValue],
262
+ usage: Sequence[UsageRecord],
263
+ ) -> None:
264
+ del usage
265
+ self._suggestion_number += 1
266
+ self._observations.append(SearchObservation(candidate, tuple(metrics)))
267
+
268
+ def fingerprint(self) -> Any:
269
+ provider_type = type(self._provider)
270
+ provider: dict[str, Any] = {
271
+ "type": f"{provider_type.__module__}.{provider_type.__qualname__}"
272
+ }
273
+ if isinstance(self._provider, Fingerprintable):
274
+ provider["fingerprint"] = self._provider.fingerprint()
275
+ return {
276
+ "space": self._space,
277
+ "provider": provider,
278
+ "max_suggestion_attempts": self._max_suggestion_attempts,
279
+ "fallback_seed": self._fallback_seed,
280
+ "fallback": self._fallback,
281
+ }
282
+
283
+
284
+ def _provider_observation(observation: SearchObservation) -> SearchObservation:
285
+ parameters = cast(dict[str, Any], redact(observation.candidate.parameters))
286
+ return SearchObservation(Candidate(parameters), observation.metrics)
@@ -0,0 +1,78 @@
1
+ """Canonical, secret-safe serialization for identity and cache keys."""
2
+
3
+ import hashlib
4
+ import json
5
+ import math
6
+ from collections.abc import Mapping, Sequence
7
+ from dataclasses import asdict, is_dataclass
8
+ from typing import Any, cast
9
+
10
+ from .domain import Secret
11
+
12
+
13
+ def canonicalize(value: Any) -> Any:
14
+ """Convert supported values to deterministic JSON-compatible structures."""
15
+
16
+ if isinstance(value, Secret):
17
+ return {"__secret__": True}
18
+ if is_dataclass(value):
19
+ return canonicalize(asdict(cast(Any, value)))
20
+ if isinstance(value, Mapping):
21
+ return {
22
+ str(key): canonicalize(item)
23
+ for key, item in sorted(value.items(), key=lambda entry: str(entry[0]))
24
+ if not isinstance(item, Secret)
25
+ }
26
+ if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
27
+ return [canonicalize(item) for item in value]
28
+ if isinstance(value, float):
29
+ if not math.isfinite(value):
30
+ raise ValueError("nonfinite numbers are not supported")
31
+ return format(value, ".15g")
32
+ if value is None or isinstance(value, (bool, int, str)):
33
+ return value
34
+ raise TypeError(f"unsupported canonical value: {type(value).__name__}")
35
+
36
+
37
+ def canonical_json(value: Any) -> bytes:
38
+ """Serialize a value using stable key ordering and schema-independent JSON."""
39
+
40
+ return json.dumps(
41
+ canonicalize(value), ensure_ascii=True, separators=(",", ":"), sort_keys=True
42
+ ).encode("utf-8")
43
+
44
+
45
+ def storage_json(value: Any) -> bytes:
46
+ """Serialize supported values reversibly while omitting mapping secrets."""
47
+
48
+ return json.dumps(
49
+ _storage_value(value), ensure_ascii=True, separators=(",", ":"), sort_keys=True
50
+ ).encode("utf-8")
51
+
52
+
53
+ def content_hash(value: Any, *, namespace: str = "autorag-tuner:v1") -> str:
54
+ """Return a stable SHA-256 digest for a canonical value."""
55
+
56
+ return hashlib.sha256(namespace.encode() + b"\0" + canonical_json(value)).hexdigest()
57
+
58
+
59
+ def _storage_value(value: Any) -> Any:
60
+ if isinstance(value, Secret):
61
+ return {"__secret__": True}
62
+ if is_dataclass(value):
63
+ return _storage_value(asdict(cast(Any, value)))
64
+ if isinstance(value, Mapping):
65
+ return {
66
+ str(key): _storage_value(item)
67
+ for key, item in sorted(value.items(), key=lambda entry: str(entry[0]))
68
+ if not isinstance(item, Secret)
69
+ }
70
+ if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
71
+ return [_storage_value(item) for item in value]
72
+ if isinstance(value, float):
73
+ if not math.isfinite(value):
74
+ raise ValueError("nonfinite numbers are not supported")
75
+ return value
76
+ if value is None or isinstance(value, (bool, int, str)):
77
+ return value
78
+ raise TypeError(f"unsupported storage value: {type(value).__name__}")
tunarag/stopping.py ADDED
@@ -0,0 +1,193 @@
1
+ """Composable stop conditions for study execution."""
2
+
3
+ import math
4
+ from dataclasses import dataclass
5
+ from typing import Protocol
6
+
7
+
8
+ @dataclass(frozen=True, slots=True)
9
+ class StopContext:
10
+ """Current study totals supplied to stop conditions."""
11
+
12
+ completed_trials: int
13
+ exact_cost: float = 0.0
14
+ estimated_cost: float = 0.0
15
+ elapsed_seconds: float = 0.0
16
+ successful_trials: int = 0
17
+ failed_trials: int = 0
18
+ objective_scores: tuple[float | None, ...] = ()
19
+
20
+ def __post_init__(self) -> None:
21
+ if (
22
+ isinstance(self.completed_trials, bool)
23
+ or not isinstance(self.completed_trials, int)
24
+ or self.completed_trials < 0
25
+ ):
26
+ raise ValueError("completed trials must not be negative")
27
+ if any(
28
+ isinstance(value, bool) or not isinstance(value, int) or value < 0
29
+ for value in (self.successful_trials, self.failed_trials)
30
+ ):
31
+ raise ValueError("trial status counts must not be negative")
32
+ if self.successful_trials + self.failed_trials > self.completed_trials:
33
+ raise ValueError("trial status counts exceed completed trials")
34
+ if self.objective_scores and len(self.objective_scores) != self.completed_trials:
35
+ raise ValueError("objective score history must match completed trials")
36
+ for name in ("exact_cost", "estimated_cost", "elapsed_seconds"):
37
+ value = getattr(self, name)
38
+ if not math.isfinite(value) or value < 0:
39
+ raise ValueError(f"{name} must be finite and nonnegative")
40
+ if any(score is not None and not math.isfinite(score) for score in self.objective_scores):
41
+ raise ValueError("objective scores must be finite or None")
42
+
43
+
44
+ class StopCondition(Protocol):
45
+ """Return whether execution should stop for the current context."""
46
+
47
+ def should_stop(self, context: StopContext) -> bool: ...
48
+
49
+
50
+ @dataclass(frozen=True, slots=True)
51
+ class MaxTrials:
52
+ """Stop after a maximum number of completed trials."""
53
+
54
+ limit: int
55
+
56
+ def __post_init__(self) -> None:
57
+ if isinstance(self.limit, bool) or not isinstance(self.limit, int) or self.limit < 1:
58
+ raise ValueError("trial limit must be at least one")
59
+
60
+ def should_stop(self, context: StopContext) -> bool:
61
+ return context.completed_trials >= self.limit
62
+
63
+ def fingerprint(self) -> object:
64
+ return {"limit": self.limit}
65
+
66
+
67
+ @dataclass(frozen=True, slots=True)
68
+ class CostBudget:
69
+ """Stop when exact cost, or exact plus estimated cost, reaches a budget."""
70
+
71
+ limit: float
72
+ count_estimated: bool = True
73
+
74
+ def __post_init__(self) -> None:
75
+ if isinstance(self.limit, bool) or not math.isfinite(self.limit) or self.limit < 0:
76
+ raise ValueError("cost budget must be finite and nonnegative")
77
+ if not isinstance(self.count_estimated, bool):
78
+ raise TypeError("count_estimated must be a boolean")
79
+
80
+ def should_stop(self, context: StopContext) -> bool:
81
+ total = context.exact_cost + (context.estimated_cost if self.count_estimated else 0.0)
82
+ return total >= self.limit
83
+
84
+ def fingerprint(self) -> object:
85
+ return {"limit": self.limit, "count_estimated": self.count_estimated}
86
+
87
+
88
+ @dataclass(frozen=True, slots=True)
89
+ class TimeBudget:
90
+ """Stop when elapsed wall-clock execution reaches a limit."""
91
+
92
+ seconds: float
93
+
94
+ def __post_init__(self) -> None:
95
+ if isinstance(self.seconds, bool) or not math.isfinite(self.seconds) or self.seconds < 0:
96
+ raise ValueError("time budget must be finite and nonnegative")
97
+
98
+ def should_stop(self, context: StopContext) -> bool:
99
+ return context.elapsed_seconds >= self.seconds
100
+
101
+ def fingerprint(self) -> object:
102
+ return {"seconds": self.seconds}
103
+
104
+
105
+ @dataclass(frozen=True, slots=True)
106
+ class TargetScore:
107
+ """Stop when any successful objective reaches a benefit-oriented target."""
108
+
109
+ target: float
110
+
111
+ def __post_init__(self) -> None:
112
+ if isinstance(self.target, bool) or not math.isfinite(self.target):
113
+ raise ValueError("target score must be finite")
114
+
115
+ def should_stop(self, context: StopContext) -> bool:
116
+ return any(score is not None and score >= self.target for score in context.objective_scores)
117
+
118
+ def fingerprint(self) -> object:
119
+ return {"target": self.target}
120
+
121
+
122
+ @dataclass(frozen=True, slots=True)
123
+ class MaxFailures:
124
+ """Stop after a total or consecutive number of failed trials."""
125
+
126
+ limit: int
127
+ consecutive: bool = False
128
+
129
+ def __post_init__(self) -> None:
130
+ if isinstance(self.limit, bool) or not isinstance(self.limit, int) or self.limit < 1:
131
+ raise ValueError("failure limit must be a positive integer")
132
+
133
+ def should_stop(self, context: StopContext) -> bool:
134
+ if not self.consecutive:
135
+ return context.failed_trials >= self.limit
136
+ failures = 0
137
+ for score in reversed(context.objective_scores):
138
+ if score is not None:
139
+ break
140
+ failures += 1
141
+ return failures >= self.limit
142
+
143
+ def fingerprint(self) -> object:
144
+ return {"limit": self.limit, "consecutive": self.consecutive}
145
+
146
+
147
+ @dataclass(frozen=True, slots=True)
148
+ class NoImprovement:
149
+ """Stop after completed trials fail to improve the best objective."""
150
+
151
+ patience: int
152
+ min_delta: float = 0.0
153
+ warmup_trials: int = 1
154
+
155
+ def __post_init__(self) -> None:
156
+ if (
157
+ isinstance(self.patience, bool)
158
+ or not isinstance(self.patience, int)
159
+ or self.patience < 1
160
+ ):
161
+ raise ValueError("no-improvement patience must be a positive integer")
162
+ if (
163
+ isinstance(self.min_delta, bool)
164
+ or not math.isfinite(self.min_delta)
165
+ or self.min_delta < 0
166
+ ):
167
+ raise ValueError("minimum improvement must be finite and nonnegative")
168
+ if (
169
+ isinstance(self.warmup_trials, bool)
170
+ or not isinstance(self.warmup_trials, int)
171
+ or self.warmup_trials < 1
172
+ ):
173
+ raise ValueError("warmup trials must be a positive integer")
174
+
175
+ def should_stop(self, context: StopContext) -> bool:
176
+ if context.completed_trials < self.warmup_trials or not context.objective_scores:
177
+ return False
178
+ best: float | None = None
179
+ since_improvement = 0
180
+ for score in context.objective_scores:
181
+ if score is not None and (best is None or score > best + self.min_delta):
182
+ best = score
183
+ since_improvement = 0
184
+ else:
185
+ since_improvement += 1
186
+ return best is not None and since_improvement >= self.patience
187
+
188
+ def fingerprint(self) -> object:
189
+ return {
190
+ "patience": self.patience,
191
+ "min_delta": self.min_delta,
192
+ "warmup_trials": self.warmup_trials,
193
+ }