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 +166 -0
- tunarag/cache.py +326 -0
- tunarag/config.py +257 -0
- tunarag/contracts.py +88 -0
- tunarag/dataset.py +481 -0
- tunarag/domain.py +83 -0
- tunarag/engine.py +979 -0
- tunarag/errors.py +179 -0
- tunarag/evaluators.py +186 -0
- tunarag/integrations/__init__.py +25 -0
- tunarag/integrations/mlflow.py +155 -0
- tunarag/integrations/runnables.py +270 -0
- tunarag/objective.py +75 -0
- tunarag/py.typed +1 -0
- tunarag/result.py +309 -0
- tunarag/retry.py +70 -0
- tunarag/search.py +286 -0
- tunarag/serialization.py +78 -0
- tunarag/stopping.py +193 -0
- tunarag/store.py +941 -0
- tunarag/synthetic.py +257 -0
- tunarag_python-0.2.1.dist-info/METADATA +1164 -0
- tunarag_python-0.2.1.dist-info/RECORD +24 -0
- tunarag_python-0.2.1.dist-info/WHEEL +4 -0
tunarag/config.py
ADDED
|
@@ -0,0 +1,257 @@
|
|
|
1
|
+
"""Typed settings and deterministic configuration resolution."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
import os
|
|
5
|
+
import sys
|
|
6
|
+
from collections.abc import Mapping
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from enum import Enum
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Any
|
|
11
|
+
|
|
12
|
+
if sys.version_info >= (3, 11):
|
|
13
|
+
import tomllib
|
|
14
|
+
else:
|
|
15
|
+
import tomli as tomllib
|
|
16
|
+
|
|
17
|
+
from .errors import ConfigurationError
|
|
18
|
+
|
|
19
|
+
_ENV_TO_FIELD = {
|
|
20
|
+
"TUNARAG_HOME": "home",
|
|
21
|
+
"TUNARAG_DB_PATH": "db_path",
|
|
22
|
+
"TUNARAG_CACHE_DIR": "cache_dir",
|
|
23
|
+
"TUNARAG_ARTIFACT_DIR": "artifact_dir",
|
|
24
|
+
"TUNARAG_MAX_TRIALS": "max_trials",
|
|
25
|
+
"TUNARAG_CONCURRENCY": "concurrency",
|
|
26
|
+
"TUNARAG_SAMPLE_CONCURRENCY": "sample_concurrency",
|
|
27
|
+
"TUNARAG_SEED": "seed",
|
|
28
|
+
"TUNARAG_MAX_RETRIES": "max_retries",
|
|
29
|
+
"TUNARAG_SAMPLE_TIMEOUT_SECONDS": "sample_timeout_seconds",
|
|
30
|
+
"TUNARAG_TRIAL_TIMEOUT_SECONDS": "trial_timeout_seconds",
|
|
31
|
+
"TUNARAG_CACHE_POLICY": "cache_policy",
|
|
32
|
+
}
|
|
33
|
+
_FIELDS = frozenset(_ENV_TO_FIELD.values())
|
|
34
|
+
_INTEGER_FIELDS = {
|
|
35
|
+
"max_trials",
|
|
36
|
+
"concurrency",
|
|
37
|
+
"sample_concurrency",
|
|
38
|
+
"seed",
|
|
39
|
+
"max_retries",
|
|
40
|
+
}
|
|
41
|
+
_PATH_FIELDS = {"home", "db_path", "cache_dir", "artifact_dir"}
|
|
42
|
+
_FLOAT_FIELDS = {"sample_timeout_seconds", "trial_timeout_seconds"}
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
class CachePolicy(str, Enum):
|
|
46
|
+
"""Allowed cache read and write behavior."""
|
|
47
|
+
|
|
48
|
+
OFF = "off"
|
|
49
|
+
READ_ONLY = "read_only"
|
|
50
|
+
WRITE_ONLY = "write_only"
|
|
51
|
+
READ_WRITE = "read_write"
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@dataclass(frozen=True, slots=True)
|
|
55
|
+
class Settings:
|
|
56
|
+
"""Resolved core settings with no provider credentials."""
|
|
57
|
+
|
|
58
|
+
home: Path
|
|
59
|
+
db_path: Path
|
|
60
|
+
cache_dir: Path
|
|
61
|
+
artifact_dir: Path
|
|
62
|
+
max_trials: int = 25
|
|
63
|
+
concurrency: int = 4
|
|
64
|
+
sample_concurrency: int = 8
|
|
65
|
+
seed: int = 0
|
|
66
|
+
max_retries: int = 2
|
|
67
|
+
sample_timeout_seconds: float = 60.0
|
|
68
|
+
trial_timeout_seconds: float = 300.0
|
|
69
|
+
cache_policy: CachePolicy = CachePolicy.READ_WRITE
|
|
70
|
+
|
|
71
|
+
def __post_init__(self) -> None:
|
|
72
|
+
_require_positive("max_trials", self.max_trials)
|
|
73
|
+
_require_positive("concurrency", self.concurrency)
|
|
74
|
+
_require_positive("sample_concurrency", self.sample_concurrency)
|
|
75
|
+
_require_integer("seed", self.seed)
|
|
76
|
+
_require_integer("max_retries", self.max_retries)
|
|
77
|
+
if self.max_retries < 0:
|
|
78
|
+
raise ConfigurationError("max_retries must not be negative", stage="config")
|
|
79
|
+
for field in _FLOAT_FIELDS:
|
|
80
|
+
_require_positive_float(field, getattr(self, field))
|
|
81
|
+
|
|
82
|
+
@classmethod
|
|
83
|
+
def resolve(
|
|
84
|
+
cls,
|
|
85
|
+
*,
|
|
86
|
+
overrides: Mapping[str, Any] | None = None,
|
|
87
|
+
config_path: str | Path | None = None,
|
|
88
|
+
environ: Mapping[str, str] | None = None,
|
|
89
|
+
) -> "Settings":
|
|
90
|
+
"""Resolve defaults, selected TOML, known environment variables, and overrides."""
|
|
91
|
+
|
|
92
|
+
values: dict[str, Any] = {
|
|
93
|
+
"max_trials": 25,
|
|
94
|
+
"concurrency": 4,
|
|
95
|
+
"sample_concurrency": 8,
|
|
96
|
+
"seed": 0,
|
|
97
|
+
"max_retries": 2,
|
|
98
|
+
"sample_timeout_seconds": 60.0,
|
|
99
|
+
"trial_timeout_seconds": 300.0,
|
|
100
|
+
"cache_policy": CachePolicy.READ_WRITE,
|
|
101
|
+
}
|
|
102
|
+
if config_path is not None:
|
|
103
|
+
values.update(_read_toml(Path(config_path)))
|
|
104
|
+
values.update(_read_environment(os.environ if environ is None else environ))
|
|
105
|
+
if overrides is not None:
|
|
106
|
+
_reject_unknown(overrides, source="overrides")
|
|
107
|
+
values.update(overrides)
|
|
108
|
+
|
|
109
|
+
home = _as_path("home", values.get("home", _default_home()))
|
|
110
|
+
values["home"] = home
|
|
111
|
+
values.setdefault("db_path", home / "experiments.db")
|
|
112
|
+
values.setdefault("cache_dir", home / "cache")
|
|
113
|
+
values.setdefault("artifact_dir", home / "artifacts")
|
|
114
|
+
|
|
115
|
+
for field in _PATH_FIELDS:
|
|
116
|
+
values[field] = _as_path(field, values[field])
|
|
117
|
+
for field in _INTEGER_FIELDS:
|
|
118
|
+
values[field] = _as_integer(field, values[field])
|
|
119
|
+
for field in _FLOAT_FIELDS:
|
|
120
|
+
values[field] = _as_float(field, values[field])
|
|
121
|
+
values["cache_policy"] = _as_cache_policy(values["cache_policy"])
|
|
122
|
+
return cls(**values)
|
|
123
|
+
|
|
124
|
+
def as_dict(self) -> dict[str, str | int | float]:
|
|
125
|
+
"""Return a serializable settings view containing no secrets."""
|
|
126
|
+
|
|
127
|
+
return {
|
|
128
|
+
"home": str(self.home),
|
|
129
|
+
"db_path": str(self.db_path),
|
|
130
|
+
"cache_dir": str(self.cache_dir),
|
|
131
|
+
"artifact_dir": str(self.artifact_dir),
|
|
132
|
+
"max_trials": self.max_trials,
|
|
133
|
+
"concurrency": self.concurrency,
|
|
134
|
+
"sample_concurrency": self.sample_concurrency,
|
|
135
|
+
"seed": self.seed,
|
|
136
|
+
"max_retries": self.max_retries,
|
|
137
|
+
"sample_timeout_seconds": self.sample_timeout_seconds,
|
|
138
|
+
"trial_timeout_seconds": self.trial_timeout_seconds,
|
|
139
|
+
"cache_policy": self.cache_policy.value,
|
|
140
|
+
}
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def _read_toml(path: Path) -> dict[str, Any]:
|
|
144
|
+
try:
|
|
145
|
+
with path.open("rb") as stream:
|
|
146
|
+
document = tomllib.load(stream)
|
|
147
|
+
except (OSError, tomllib.TOMLDecodeError) as error:
|
|
148
|
+
raise ConfigurationError(
|
|
149
|
+
f"could not read TunaRAG configuration: {path}",
|
|
150
|
+
stage="config",
|
|
151
|
+
details={"path": str(path), "cause": type(error).__name__},
|
|
152
|
+
) from error
|
|
153
|
+
section = document.get("tunarag", {})
|
|
154
|
+
if not isinstance(section, Mapping):
|
|
155
|
+
raise ConfigurationError("the [tunarag] TOML section must be a table", stage="config")
|
|
156
|
+
_reject_unknown(section, source=str(path))
|
|
157
|
+
return dict(section)
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
def _read_environment(environ: Mapping[str, str]) -> dict[str, Any]:
|
|
161
|
+
values: dict[str, Any] = {}
|
|
162
|
+
for variable, field in _ENV_TO_FIELD.items():
|
|
163
|
+
if variable not in environ:
|
|
164
|
+
continue
|
|
165
|
+
raw = environ[variable]
|
|
166
|
+
if not raw.strip():
|
|
167
|
+
raise ConfigurationError(
|
|
168
|
+
f"{variable} must not be empty",
|
|
169
|
+
stage="config",
|
|
170
|
+
details={"variable": variable},
|
|
171
|
+
)
|
|
172
|
+
values[field] = raw
|
|
173
|
+
return values
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
def _reject_unknown(values: Mapping[str, Any], *, source: str) -> None:
|
|
177
|
+
unknown = sorted(str(key) for key in values if key not in _FIELDS)
|
|
178
|
+
if unknown:
|
|
179
|
+
raise ConfigurationError(
|
|
180
|
+
f"unknown TunaRAG configuration keys in {source}",
|
|
181
|
+
stage="config",
|
|
182
|
+
details={"keys": unknown, "source": source},
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def _as_path(field: str, value: Any) -> Path:
|
|
187
|
+
if not isinstance(value, (str, os.PathLike)):
|
|
188
|
+
raise ConfigurationError(f"{field} must be a path", stage="config")
|
|
189
|
+
if isinstance(value, str) and not value.strip():
|
|
190
|
+
raise ConfigurationError(f"{field} must not be empty", stage="config")
|
|
191
|
+
return Path(value).expanduser()
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def _as_integer(field: str, value: Any) -> int:
|
|
195
|
+
if isinstance(value, bool):
|
|
196
|
+
raise ConfigurationError(f"{field} must be an integer", stage="config")
|
|
197
|
+
try:
|
|
198
|
+
converted = int(value)
|
|
199
|
+
except (TypeError, ValueError) as error:
|
|
200
|
+
raise ConfigurationError(f"{field} must be an integer", stage="config") from error
|
|
201
|
+
if isinstance(value, float) and not value.is_integer():
|
|
202
|
+
raise ConfigurationError(f"{field} must be an integer", stage="config")
|
|
203
|
+
if isinstance(value, str) and str(converted) != value.strip():
|
|
204
|
+
raise ConfigurationError(f"{field} must be an integer", stage="config")
|
|
205
|
+
return converted
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def _as_cache_policy(value: Any) -> CachePolicy:
|
|
209
|
+
if isinstance(value, CachePolicy):
|
|
210
|
+
return value
|
|
211
|
+
try:
|
|
212
|
+
return CachePolicy(value)
|
|
213
|
+
except (TypeError, ValueError) as error:
|
|
214
|
+
choices = ", ".join(policy.value for policy in CachePolicy)
|
|
215
|
+
raise ConfigurationError(
|
|
216
|
+
f"cache_policy must be one of: {choices}", stage="config"
|
|
217
|
+
) from error
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
def _as_float(field: str, value: Any) -> float:
|
|
221
|
+
if isinstance(value, bool):
|
|
222
|
+
raise ConfigurationError(f"{field} must be a number", stage="config")
|
|
223
|
+
try:
|
|
224
|
+
converted = float(value)
|
|
225
|
+
except (TypeError, ValueError) as error:
|
|
226
|
+
raise ConfigurationError(f"{field} must be a number", stage="config") from error
|
|
227
|
+
if not math.isfinite(converted):
|
|
228
|
+
raise ConfigurationError(f"{field} must be finite", stage="config")
|
|
229
|
+
return converted
|
|
230
|
+
|
|
231
|
+
|
|
232
|
+
def _require_integer(field: str, value: int) -> None:
|
|
233
|
+
if not isinstance(value, int) or isinstance(value, bool):
|
|
234
|
+
raise ConfigurationError(f"{field} must be an integer", stage="config")
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def _require_positive(field: str, value: int) -> None:
|
|
238
|
+
_require_integer(field, value)
|
|
239
|
+
if value < 1:
|
|
240
|
+
raise ConfigurationError(f"{field} must be positive", stage="config")
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def _require_positive_float(field: str, value: float) -> None:
|
|
244
|
+
if not isinstance(value, (int, float)) or isinstance(value, bool):
|
|
245
|
+
raise ConfigurationError(f"{field} must be a number", stage="config")
|
|
246
|
+
if not math.isfinite(value) or value <= 0:
|
|
247
|
+
raise ConfigurationError(f"{field} must be a positive finite number", stage="config")
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
def _default_home() -> Path:
|
|
251
|
+
if sys.platform == "win32":
|
|
252
|
+
root = os.environ.get("LOCALAPPDATA")
|
|
253
|
+
return Path(root) / "TunaRAG" if root else Path.home() / "AppData" / "Local" / "TunaRAG"
|
|
254
|
+
if sys.platform == "darwin":
|
|
255
|
+
return Path.home() / "Library" / "Application Support" / "TunaRAG"
|
|
256
|
+
root = os.environ.get("XDG_DATA_HOME")
|
|
257
|
+
return Path(root) / "tunarag" if root else Path.home() / ".local" / "share" / "tunarag"
|
tunarag/contracts.py
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
"""Protocols implemented by user-owned RAG systems and package extensions."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import AsyncIterator, Sequence
|
|
6
|
+
from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable
|
|
7
|
+
|
|
8
|
+
from .cache import CacheEntry, CacheKey, CacheLookup, CachePrivacy
|
|
9
|
+
from .domain import Candidate, MetricValue, UsageRecord
|
|
10
|
+
|
|
11
|
+
if TYPE_CHECKING:
|
|
12
|
+
from .store import StudyEvent
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class RAGAdapter(Protocol):
|
|
16
|
+
"""Execute a user's RAG system for one candidate and example."""
|
|
17
|
+
|
|
18
|
+
async def run(
|
|
19
|
+
self, candidate: Candidate, example: Any
|
|
20
|
+
) -> tuple[Any, Sequence[UsageRecord]]: ...
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class Dataset(Protocol):
|
|
24
|
+
"""Provide versioned evaluation examples."""
|
|
25
|
+
|
|
26
|
+
@property
|
|
27
|
+
def version(self) -> str: ...
|
|
28
|
+
|
|
29
|
+
def __aiter__(self) -> AsyncIterator[Any]: ...
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class Evaluator(Protocol):
|
|
33
|
+
"""Convert an adapter output and example into metric observations."""
|
|
34
|
+
|
|
35
|
+
async def evaluate(self, output: Any, example: Any) -> Sequence[MetricValue]: ...
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class SearchStrategy(Protocol):
|
|
39
|
+
"""Suggest candidates and observe completed trials."""
|
|
40
|
+
|
|
41
|
+
async def suggest(self) -> Candidate: ...
|
|
42
|
+
|
|
43
|
+
async def observe(
|
|
44
|
+
self,
|
|
45
|
+
candidate: Candidate,
|
|
46
|
+
metrics: Sequence[MetricValue],
|
|
47
|
+
usage: Sequence[UsageRecord],
|
|
48
|
+
) -> None: ...
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@runtime_checkable
|
|
52
|
+
class ReplayableSearchStrategy(Protocol):
|
|
53
|
+
"""Restore strategy state from one durably reserved trial."""
|
|
54
|
+
|
|
55
|
+
async def replay(
|
|
56
|
+
self,
|
|
57
|
+
candidate: Candidate,
|
|
58
|
+
metrics: Sequence[MetricValue],
|
|
59
|
+
usage: Sequence[UsageRecord],
|
|
60
|
+
) -> None: ...
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
@runtime_checkable
|
|
64
|
+
class Fingerprintable(Protocol):
|
|
65
|
+
"""Expose secret-safe semantic configuration for identity checks."""
|
|
66
|
+
|
|
67
|
+
def fingerprint(self) -> Any: ...
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
class EventCallback(Protocol):
|
|
71
|
+
"""Consume a durable event after its transaction commits."""
|
|
72
|
+
|
|
73
|
+
async def on_event(self, event: StudyEvent) -> None: ...
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class Cache(Protocol):
|
|
77
|
+
"""Persist and retrieve versioned semantic results."""
|
|
78
|
+
|
|
79
|
+
def get(self, key: CacheKey) -> CacheLookup: ...
|
|
80
|
+
|
|
81
|
+
def put(
|
|
82
|
+
self,
|
|
83
|
+
key: CacheKey,
|
|
84
|
+
payload: bytes | bytearray | memoryview,
|
|
85
|
+
*,
|
|
86
|
+
privacy: CachePrivacy = CachePrivacy.STANDARD,
|
|
87
|
+
ttl_seconds: float | None = None,
|
|
88
|
+
) -> CacheEntry: ...
|