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/errors.py ADDED
@@ -0,0 +1,179 @@
1
+ """Structured public errors with secret-safe diagnostic details."""
2
+
3
+ from collections.abc import Mapping, Sequence
4
+ from dataclasses import fields, is_dataclass
5
+ from enum import Enum
6
+ from types import MappingProxyType
7
+ from typing import Any, cast
8
+
9
+ from .domain import Secret
10
+
11
+ _SENSITIVE_KEYS = {
12
+ "access_token",
13
+ "api_key",
14
+ "authorization",
15
+ "credential",
16
+ "credentials",
17
+ "password",
18
+ "refresh_token",
19
+ "secret",
20
+ "token",
21
+ }
22
+ _REDACTED = "***redacted***"
23
+
24
+
25
+ class ErrorCode(str, Enum):
26
+ """Stable machine-readable error identifiers."""
27
+
28
+ CONFIGURATION = "configuration.invalid"
29
+ SEARCH_SPACE = "search_space.invalid"
30
+ DATASET = "dataset.invalid"
31
+ ADAPTER = "adapter.failed"
32
+ EVALUATION = "evaluation.failed"
33
+ STORE = "store.failed"
34
+ RESUME = "resume.incompatible"
35
+ BUDGET = "budget.exhausted"
36
+ OPTIONAL_DEPENDENCY = "dependency.optional_missing"
37
+ SYNC_IN_ASYNC = "runtime.sync_in_async"
38
+
39
+
40
+ class TunaRAGError(Exception):
41
+ """Base class for public package errors."""
42
+
43
+ code = ErrorCode.CONFIGURATION
44
+
45
+ def __init__(
46
+ self,
47
+ message: str,
48
+ *,
49
+ stage: str | None = None,
50
+ retryable: bool = False,
51
+ details: Mapping[str, Any] | None = None,
52
+ ) -> None:
53
+ if not message.strip():
54
+ raise ValueError("error message must not be empty")
55
+ if stage is not None and not stage.strip():
56
+ raise ValueError("error stage must not be empty")
57
+ self.message = message
58
+ self.stage = stage
59
+ self.retryable = retryable
60
+ safe_details = cast(dict[str, Any], redact(details or {}))
61
+ self.details: Mapping[str, Any] = MappingProxyType(safe_details)
62
+ super().__init__(message)
63
+
64
+ def as_dict(self) -> dict[str, Any]:
65
+ """Return a JSON-compatible diagnostic representation."""
66
+
67
+ return {
68
+ "code": self.code.value,
69
+ "message": self.message,
70
+ "stage": self.stage,
71
+ "retryable": self.retryable,
72
+ "details": dict(self.details),
73
+ }
74
+
75
+
76
+ class ConfigurationError(TunaRAGError):
77
+ """Configuration could not be parsed or validated."""
78
+
79
+ code = ErrorCode.CONFIGURATION
80
+
81
+
82
+ class SearchSpaceError(TunaRAGError):
83
+ """A search-space definition or candidate is invalid."""
84
+
85
+ code = ErrorCode.SEARCH_SPACE
86
+
87
+
88
+ class DatasetError(TunaRAGError):
89
+ """An evaluation dataset could not be loaded or validated."""
90
+
91
+ code = ErrorCode.DATASET
92
+
93
+ def __init__(self, message: str, *, details: Mapping[str, Any] | None = None) -> None:
94
+ super().__init__(message, stage="dataset", details=details)
95
+
96
+
97
+ class AdapterError(TunaRAGError):
98
+ """The user-owned RAG adapter failed."""
99
+
100
+ code = ErrorCode.ADAPTER
101
+
102
+
103
+ class EvaluationError(TunaRAGError):
104
+ """An evaluator failed or returned invalid observations."""
105
+
106
+ code = ErrorCode.EVALUATION
107
+
108
+
109
+ class StoreError(TunaRAGError):
110
+ """Durable state could not be read or written."""
111
+
112
+ code = ErrorCode.STORE
113
+
114
+
115
+ class ResumeError(TunaRAGError):
116
+ """A study cannot resume under the supplied semantics."""
117
+
118
+ code = ErrorCode.RESUME
119
+
120
+
121
+ class BudgetError(TunaRAGError):
122
+ """A configured scheduling budget has been exhausted."""
123
+
124
+ code = ErrorCode.BUDGET
125
+
126
+
127
+ class MissingOptionalDependencyError(TunaRAGError):
128
+ """A requested integration extra is not installed."""
129
+
130
+ code = ErrorCode.OPTIONAL_DEPENDENCY
131
+
132
+ def __init__(self, feature: str, extra: str) -> None:
133
+ if not feature.strip() or not extra.strip():
134
+ raise ValueError("feature and extra must not be empty")
135
+ super().__init__(
136
+ f"{feature} requires the optional '{extra}' extra",
137
+ stage="dependency",
138
+ details={"feature": feature, "extra": extra},
139
+ )
140
+
141
+
142
+ class SyncInAsyncContextError(TunaRAGError):
143
+ """A synchronous entry point was called from an active event loop."""
144
+
145
+ code = ErrorCode.SYNC_IN_ASYNC
146
+
147
+ def __init__(self) -> None:
148
+ super().__init__(
149
+ "the synchronous API cannot run inside an active event loop; use the async API",
150
+ stage="runtime",
151
+ )
152
+
153
+
154
+ def redact(value: Any, *, key: str | None = None) -> Any:
155
+ """Return a recursively redacted, JSON-compatible diagnostic value."""
156
+
157
+ if key is not None and _is_sensitive_key(key):
158
+ return _REDACTED
159
+ if isinstance(value, Secret):
160
+ return _REDACTED
161
+ if is_dataclass(value):
162
+ return {
163
+ field.name: redact(getattr(value, field.name), key=field.name)
164
+ for field in fields(cast(Any, value))
165
+ }
166
+ if isinstance(value, Mapping):
167
+ return {str(item_key): redact(item, key=str(item_key)) for item_key, item in value.items()}
168
+ if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)):
169
+ return [redact(item) for item in value]
170
+ if value is None or isinstance(value, (bool, int, float, str)):
171
+ return value
172
+ return f"<{type(value).__name__}>"
173
+
174
+
175
+ def _is_sensitive_key(key: str) -> bool:
176
+ normalized = key.lower().replace("-", "_")
177
+ return normalized in _SENSITIVE_KEYS or normalized.endswith(
178
+ ("_api_key", "_credential", "_password", "_secret", "_token")
179
+ )
tunarag/evaluators.py ADDED
@@ -0,0 +1,186 @@
1
+ """Built-in evaluator helpers and optional third-party integrations."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import inspect
6
+ import math
7
+ from collections.abc import Awaitable, Callable, Mapping, Sequence
8
+ from importlib import import_module
9
+ from numbers import Real
10
+ from typing import Any
11
+
12
+ from .contracts import Fingerprintable
13
+ from .domain import MetricValue, ValueStatus
14
+ from .errors import EvaluationError, MissingOptionalDependencyError
15
+
16
+ MetricScorer = Callable[[Any, Any], float | None | Awaitable[float | None]]
17
+ RagasSampleMapper = Callable[[Any, Any], Mapping[str, Any]]
18
+
19
+
20
+ class MetricEvaluator:
21
+ """Adapt one synchronous or asynchronous numeric scorer to an evaluator."""
22
+
23
+ def __init__(
24
+ self,
25
+ name: str,
26
+ scorer: MetricScorer,
27
+ *,
28
+ status: ValueStatus = ValueStatus.EXACT,
29
+ ) -> None:
30
+ if not name.strip():
31
+ raise ValueError("metric name must not be empty")
32
+ if name == "tunarag.objective":
33
+ raise ValueError("metric name is reserved: tunarag.objective")
34
+ if not callable(scorer):
35
+ raise TypeError("metric scorer must be callable")
36
+ self._name = name
37
+ self._scorer = scorer
38
+ self._status = status
39
+
40
+ async def evaluate(self, output: Any, example: Any) -> Sequence[MetricValue]:
41
+ """Score one adapter output and evaluation example."""
42
+
43
+ result = self._scorer(output, example)
44
+ value = await result if inspect.isawaitable(result) else result
45
+ normalized = _numeric_score(value, metric=self._name)
46
+ if normalized is None:
47
+ return (MetricValue(self._name, None, ValueStatus.UNAVAILABLE, 0.0),)
48
+ return (MetricValue(self._name, normalized, self._status),)
49
+
50
+ def fingerprint(self) -> Any:
51
+ """Return stable configuration used by study and cache identity."""
52
+
53
+ return {
54
+ "name": self._name,
55
+ "status": self._status.value,
56
+ "scorer": _callable_identity(self._scorer),
57
+ }
58
+
59
+
60
+ class RagasEvaluator:
61
+ """Evaluate one sample with optional RAGAS metrics without hard coupling."""
62
+
63
+ def __init__(
64
+ self,
65
+ metrics: Sequence[Any],
66
+ sample_mapper: RagasSampleMapper,
67
+ *,
68
+ metric_prefix: str = "ragas.",
69
+ llm: Any = None,
70
+ embeddings: Any = None,
71
+ raise_exceptions: bool = False,
72
+ ) -> None:
73
+ if not metrics:
74
+ raise ValueError("RAGAS evaluator requires at least one metric")
75
+ if not callable(sample_mapper):
76
+ raise TypeError("RAGAS sample mapper must be callable")
77
+ if metric_prefix == "tunarag.objective" or metric_prefix.startswith("tunarag.objective."):
78
+ raise ValueError("RAGAS metric prefix uses a reserved name")
79
+ self._metrics = tuple(metrics)
80
+ self._sample_mapper = sample_mapper
81
+ self._metric_prefix = metric_prefix
82
+ self._llm = llm
83
+ self._embeddings = embeddings
84
+ self._raise_exceptions = raise_exceptions
85
+
86
+ async def evaluate(self, output: Any, example: Any) -> Sequence[MetricValue]:
87
+ """Map and evaluate exactly one sample through RAGAS's async API."""
88
+
89
+ try:
90
+ sample_data = self._sample_mapper(output, example)
91
+ except Exception as error:
92
+ raise EvaluationError(
93
+ "RAGAS sample mapping failed",
94
+ stage="evaluation",
95
+ details={"cause": type(error).__name__},
96
+ ) from error
97
+ if not isinstance(sample_data, Mapping):
98
+ raise EvaluationError("RAGAS sample mapper must return a mapping", stage="evaluation")
99
+ ragas, dataset_schema = _load_ragas()
100
+ try:
101
+ sample = dataset_schema.SingleTurnSample(**dict(sample_data))
102
+ dataset = ragas.EvaluationDataset(samples=[sample])
103
+ result = await ragas.aevaluate(
104
+ dataset=dataset,
105
+ metrics=list(self._metrics),
106
+ llm=self._llm,
107
+ embeddings=self._embeddings,
108
+ raise_exceptions=self._raise_exceptions,
109
+ show_progress=False,
110
+ )
111
+ scores = result.scores
112
+ except Exception as error:
113
+ raise EvaluationError(
114
+ "RAGAS evaluation failed",
115
+ stage="evaluation",
116
+ details={"cause": type(error).__name__},
117
+ ) from error
118
+ if not isinstance(scores, Sequence) or len(scores) != 1:
119
+ raise EvaluationError("RAGAS returned an invalid per-sample result", stage="evaluation")
120
+ row = scores[0]
121
+ if not isinstance(row, Mapping) or not row:
122
+ raise EvaluationError("RAGAS returned no metric scores", stage="evaluation")
123
+ observations: list[MetricValue] = []
124
+ names: set[str] = set()
125
+ for raw_name, raw_value in row.items():
126
+ name = f"{self._metric_prefix}{raw_name}"
127
+ if not name.strip() or name == "tunarag.objective":
128
+ raise EvaluationError("RAGAS returned an invalid metric name", stage="evaluation")
129
+ if name in names:
130
+ raise EvaluationError("RAGAS returned duplicate metric names", stage="evaluation")
131
+ names.add(name)
132
+ value = _numeric_score(raw_value, metric=name)
133
+ if value is None:
134
+ observations.append(MetricValue(name, None, ValueStatus.UNAVAILABLE, 0.0))
135
+ else:
136
+ observations.append(MetricValue(name, value))
137
+ return tuple(observations)
138
+
139
+ def fingerprint(self) -> Any:
140
+ """Return secret-safe semantic integration configuration."""
141
+
142
+ return {
143
+ "metrics": [_component_identity(metric) for metric in self._metrics],
144
+ "sample_mapper": _callable_identity(self._sample_mapper),
145
+ "metric_prefix": self._metric_prefix,
146
+ "llm": _component_identity(self._llm),
147
+ "embeddings": _component_identity(self._embeddings),
148
+ "raise_exceptions": self._raise_exceptions,
149
+ }
150
+
151
+
152
+ def _load_ragas() -> tuple[Any, Any]:
153
+ try:
154
+ return import_module("ragas"), import_module("ragas.dataset_schema")
155
+ except ImportError as error:
156
+ raise MissingOptionalDependencyError("RAGAS evaluation", "ragas") from error
157
+
158
+
159
+ def _numeric_score(value: Any, *, metric: str) -> float | None:
160
+ if value is None:
161
+ return None
162
+ if isinstance(value, bool) or not isinstance(value, Real):
163
+ raise EvaluationError(
164
+ "evaluator score must be numeric or None",
165
+ stage="evaluation",
166
+ details={"metric": metric, "value_type": type(value).__name__},
167
+ )
168
+ normalized = float(value)
169
+ return normalized if math.isfinite(normalized) else None
170
+
171
+
172
+ def _callable_identity(function: Callable[..., Any]) -> str:
173
+ module = getattr(function, "__module__", type(function).__module__)
174
+ name = getattr(function, "__qualname__", type(function).__qualname__)
175
+ return f"{module}.{name}"
176
+
177
+
178
+ def _component_identity(component: Any) -> Any:
179
+ if component is None:
180
+ return None
181
+ identity: dict[str, Any] = {
182
+ "type": f"{type(component).__module__}.{type(component).__qualname__}"
183
+ }
184
+ if isinstance(component, Fingerprintable):
185
+ identity["fingerprint"] = component.fingerprint()
186
+ return identity
@@ -0,0 +1,25 @@
1
+ """Optional integration adapters."""
2
+
3
+ from .mlflow import MLflowCallback
4
+ from .runnables import (
5
+ AsyncRunnable,
6
+ LangChainAdapter,
7
+ LangGraphAdapter,
8
+ RunnableConfigMapper,
9
+ RunnableFactory,
10
+ RunnableInputMapper,
11
+ RunnableOutputMapper,
12
+ RunnableUsageMapper,
13
+ )
14
+
15
+ __all__ = [
16
+ "AsyncRunnable",
17
+ "LangChainAdapter",
18
+ "LangGraphAdapter",
19
+ "MLflowCallback",
20
+ "RunnableConfigMapper",
21
+ "RunnableFactory",
22
+ "RunnableInputMapper",
23
+ "RunnableOutputMapper",
24
+ "RunnableUsageMapper",
25
+ ]
@@ -0,0 +1,155 @@
1
+ """Optional post-commit MLflow tracking mirror."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import hashlib
7
+ import json
8
+ import re
9
+ from collections.abc import Mapping
10
+ from dataclasses import dataclass
11
+ from importlib import import_module
12
+ from typing import Any, Protocol
13
+
14
+ from ..errors import MissingOptionalDependencyError, redact
15
+ from ..store import SQLiteStore, StudyEvent, TrialStatus
16
+
17
+
18
+ class MLflowClient(Protocol):
19
+ """Subset of the low-level MLflow client used by the mirror."""
20
+
21
+ def create_run(self, experiment_id: str, *, tags: Mapping[str, str]) -> Any: ...
22
+
23
+ def log_metric(
24
+ self, run_id: str, key: str, value: float, *, step: int | None = None
25
+ ) -> Any: ...
26
+
27
+ def log_param(self, run_id: str, key: str, value: Any) -> Any: ...
28
+
29
+ def set_tag(self, run_id: str, key: str, value: Any) -> None: ...
30
+
31
+ def set_terminated(self, run_id: str, *, status: str | None = None) -> None: ...
32
+
33
+
34
+ @dataclass(frozen=True, slots=True)
35
+ class _TrialMirrorData:
36
+ sequence: int
37
+ status: TrialStatus
38
+ candidate: Mapping[str, Any]
39
+ metrics: Mapping[str, float]
40
+
41
+
42
+ class MLflowCallback:
43
+ """Mirror committed study events into one low-level MLflow run."""
44
+
45
+ def __init__(
46
+ self,
47
+ store: SQLiteStore,
48
+ *,
49
+ experiment_id: str = "0",
50
+ tracking_uri: str | None = None,
51
+ tags: Mapping[str, str] | None = None,
52
+ client: MLflowClient | None = None,
53
+ ) -> None:
54
+ if not experiment_id.strip():
55
+ raise ValueError("MLflow experiment id must not be empty")
56
+ self._store = store
57
+ self._experiment_id = experiment_id
58
+ self._tracking_uri = tracking_uri
59
+ safe_tags = redact(tags or {})
60
+ if not isinstance(safe_tags, Mapping):
61
+ raise TypeError("MLflow tags must be a mapping")
62
+ self._tags = {str(key): str(value) for key, value in safe_tags.items()}
63
+ self._client = client
64
+ self._run_ids: dict[str, str] = {}
65
+
66
+ async def on_event(self, event: StudyEvent) -> None:
67
+ """Mirror one durable event after gathering committed local state."""
68
+
69
+ trial_data = self._trial_data(event)
70
+ await asyncio.to_thread(self._mirror, event, trial_data)
71
+
72
+ def _trial_data(self, event: StudyEvent) -> _TrialMirrorData | None:
73
+ if event.type not in {"trial.succeeded", "trial.failed", "trial.cancelled"}:
74
+ return None
75
+ trial_id = event.payload.get("trial_id")
76
+ if not isinstance(trial_id, str):
77
+ return None
78
+ trial = self._store.get_trial(trial_id)
79
+ candidate = json.loads(trial.candidate_json)
80
+ if not isinstance(candidate, Mapping):
81
+ candidate = {}
82
+ safe_candidate = redact(candidate)
83
+ if not isinstance(safe_candidate, Mapping):
84
+ safe_candidate = {}
85
+ metrics = {
86
+ metric.name: float(metric.value)
87
+ for metric in self._store.list_metrics(trial_id)
88
+ if metric.value is not None
89
+ }
90
+ return _TrialMirrorData(trial.sequence, trial.status, safe_candidate, metrics)
91
+
92
+ def _mirror(self, event: StudyEvent, trial: _TrialMirrorData | None) -> None:
93
+ client = self._client_or_import()
94
+ run_id = self._run_ids.get(event.study_id)
95
+ if run_id is None:
96
+ tags = {
97
+ **self._tags,
98
+ "tunarag.study_id": event.study_id,
99
+ "tunarag.mirror": "event-callback",
100
+ }
101
+ run = client.create_run(self._experiment_id, tags=tags)
102
+ run_id = str(run.info.run_id)
103
+ self._run_ids[event.study_id] = run_id
104
+ client.set_tag(run_id, "tunarag.last_event", event.type)
105
+ client.set_tag(run_id, "tunarag.last_event_sequence", event.sequence)
106
+ if trial is not None:
107
+ self._mirror_trial(client, run_id, trial)
108
+ terminal_status = {
109
+ "study.completed": "FINISHED",
110
+ "study.failed": "FAILED",
111
+ "study.cancelled": "KILLED",
112
+ }.get(event.type)
113
+ if terminal_status is not None:
114
+ client.set_terminated(run_id, status=terminal_status)
115
+
116
+ def _mirror_trial(self, client: MLflowClient, run_id: str, trial: _TrialMirrorData) -> None:
117
+ prefix = f"trial.{trial.sequence}"
118
+ client.set_tag(run_id, f"{prefix}.status", trial.status.value)
119
+ for key, value in sorted(trial.candidate.items()):
120
+ client.log_param(
121
+ run_id, _tracking_key(f"{prefix}.candidate.{key}"), _parameter_value(value)
122
+ )
123
+ for name, value in sorted(trial.metrics.items()):
124
+ client.log_metric(run_id, _tracking_key(name), value, step=trial.sequence)
125
+ client.log_metric(
126
+ run_id,
127
+ "tunarag.trial_failed",
128
+ 1.0 if trial.status is TrialStatus.FAILED else 0.0,
129
+ step=trial.sequence,
130
+ )
131
+
132
+ def _client_or_import(self) -> MLflowClient:
133
+ if self._client is not None:
134
+ return self._client
135
+ try:
136
+ module = import_module("mlflow")
137
+ except ImportError as error:
138
+ raise MissingOptionalDependencyError("MLflow tracking", "mlflow") from error
139
+ self._client = module.MlflowClient(tracking_uri=self._tracking_uri)
140
+ return self._client
141
+
142
+
143
+ def _parameter_value(value: Any) -> str:
144
+ safe = redact(value)
145
+ if safe is None or isinstance(safe, (bool, int, float, str)):
146
+ return str(safe)
147
+ return json.dumps(safe, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
148
+
149
+
150
+ def _tracking_key(value: str) -> str:
151
+ normalized = re.sub(r"[^A-Za-z0-9_.\-/ ]", "_", value)
152
+ if normalized == value and len(normalized) <= 250:
153
+ return normalized
154
+ digest = hashlib.sha256(value.encode("utf-8")).hexdigest()[:12]
155
+ return f"{normalized[:237]}.{digest}"