|
14 | 14 |
|
15 | 15 | from __future__ import annotations |
16 | 16 |
|
| 17 | +from google.adk.agents.common_configs import CodeConfig |
17 | 18 | from google.adk.errors.not_found_error import NotFoundError |
| 19 | +from google.adk.evaluation.custom_metric_evaluator import _CustomMetricEvaluator |
| 20 | +from google.adk.evaluation.eval_config import CustomMetricConfig |
| 21 | +from google.adk.evaluation.eval_config import EvalConfig |
18 | 22 | from google.adk.evaluation.eval_metrics import EvalMetric |
19 | 23 | from google.adk.evaluation.eval_metrics import Interval |
20 | 24 | from google.adk.evaluation.eval_metrics import MetricInfo |
21 | 25 | from google.adk.evaluation.eval_metrics import MetricValueInfo |
22 | 26 | from google.adk.evaluation.eval_metrics import PrebuiltMetrics |
23 | 27 | from google.adk.evaluation.evaluator import Evaluator |
| 28 | +from google.adk.evaluation.metric_evaluator_registry import DEFAULT_METRIC_EVALUATOR_REGISTRY |
24 | 29 | from google.adk.evaluation.metric_evaluator_registry import FinalResponseMatchV2EvaluatorMetricInfoProvider |
25 | 30 | from google.adk.evaluation.metric_evaluator_registry import HallucinationsV1EvaluatorMetricInfoProvider |
26 | 31 | from google.adk.evaluation.metric_evaluator_registry import MetricEvaluatorRegistry |
27 | 32 | from google.adk.evaluation.metric_evaluator_registry import PerTurnUserSimulatorQualityV1MetricInfoProvider |
| 33 | +from google.adk.evaluation.metric_evaluator_registry import register_custom_metrics_from_config |
28 | 34 | from google.adk.evaluation.metric_evaluator_registry import ResponseEvaluatorMetricInfoProvider |
29 | 35 | from google.adk.evaluation.metric_evaluator_registry import RubricBasedFinalResponseQualityV1EvaluatorMetricInfoProvider |
30 | 36 | from google.adk.evaluation.metric_evaluator_registry import RubricBasedMultiTurnTrajectoryMetricInfoProvider |
@@ -120,6 +126,112 @@ def test_get_evaluator_not_found(self, registry): |
120 | 126 | registry.get_evaluator(eval_metric) |
121 | 127 |
|
122 | 128 |
|
| 129 | +class TestRegisterCustomMetricsFromConfig: |
| 130 | + """Test cases for register_custom_metrics_from_config.""" |
| 131 | + |
| 132 | + _CUSTOM_METRIC_NAME = "custom_metric_for_registry_test" |
| 133 | + |
| 134 | + @pytest.fixture |
| 135 | + def registry(self): |
| 136 | + registry = MetricEvaluatorRegistry() |
| 137 | + yield registry |
| 138 | + # The registry dict is shared class-level state; remove what we added. |
| 139 | + registry._registry.pop(self._CUSTOM_METRIC_NAME, None) |
| 140 | + |
| 141 | + def _registered_metric_info(self, registry, metric_name): |
| 142 | + return next( |
| 143 | + metric_info |
| 144 | + for metric_info in registry.get_registered_metrics() |
| 145 | + if metric_info.metric_name == metric_name |
| 146 | + ) |
| 147 | + |
| 148 | + def test_registers_custom_metric_with_provided_metric_info(self, registry): |
| 149 | + metric_info = MetricInfo( |
| 150 | + metric_name="name_to_be_overridden", |
| 151 | + description="Custom metric description", |
| 152 | + metric_value_info=MetricValueInfo( |
| 153 | + interval=Interval(min_value=0.0, max_value=5.0) |
| 154 | + ), |
| 155 | + ) |
| 156 | + eval_config = EvalConfig( |
| 157 | + custom_metrics={ |
| 158 | + self._CUSTOM_METRIC_NAME: CustomMetricConfig( |
| 159 | + code_config=CodeConfig(name="math.sqrt"), |
| 160 | + metric_info=metric_info, |
| 161 | + ) |
| 162 | + } |
| 163 | + ) |
| 164 | + |
| 165 | + result = register_custom_metrics_from_config(eval_config, registry) |
| 166 | + |
| 167 | + assert result is registry |
| 168 | + registered_info = self._registered_metric_info( |
| 169 | + registry, self._CUSTOM_METRIC_NAME |
| 170 | + ) |
| 171 | + assert registered_info.metric_value_info.interval.max_value == 5.0 |
| 172 | + assert all( |
| 173 | + metric_info.metric_name != "name_to_be_overridden" |
| 174 | + for metric_info in registry.get_registered_metrics() |
| 175 | + ) |
| 176 | + evaluator = registry.get_evaluator( |
| 177 | + EvalMetric( |
| 178 | + metric_name=self._CUSTOM_METRIC_NAME, |
| 179 | + threshold=0.5, |
| 180 | + custom_function_path="math.sqrt", |
| 181 | + ) |
| 182 | + ) |
| 183 | + assert isinstance(evaluator, _CustomMetricEvaluator) |
| 184 | + |
| 185 | + def test_registers_custom_metric_with_default_metric_info(self, registry): |
| 186 | + eval_config = EvalConfig( |
| 187 | + custom_metrics={ |
| 188 | + self._CUSTOM_METRIC_NAME: CustomMetricConfig( |
| 189 | + code_config=CodeConfig(name="math.sqrt"), |
| 190 | + description="A custom metric", |
| 191 | + ) |
| 192 | + } |
| 193 | + ) |
| 194 | + |
| 195 | + register_custom_metrics_from_config(eval_config, registry) |
| 196 | + |
| 197 | + registered_info = self._registered_metric_info( |
| 198 | + registry, self._CUSTOM_METRIC_NAME |
| 199 | + ) |
| 200 | + assert registered_info.description == "A custom metric" |
| 201 | + assert registered_info.metric_value_info.interval.min_value == 0.0 |
| 202 | + assert registered_info.metric_value_info.interval.max_value == 1.0 |
| 203 | + |
| 204 | + def test_no_custom_metrics_is_a_no_op(self, registry): |
| 205 | + registered_before = registry.get_registered_metrics() |
| 206 | + |
| 207 | + result = register_custom_metrics_from_config(EvalConfig(), registry) |
| 208 | + |
| 209 | + assert result is registry |
| 210 | + assert registry.get_registered_metrics() == registered_before |
| 211 | + |
| 212 | + def test_defaults_to_the_default_registry(self): |
| 213 | + eval_config = EvalConfig( |
| 214 | + custom_metrics={ |
| 215 | + self._CUSTOM_METRIC_NAME: CustomMetricConfig( |
| 216 | + code_config=CodeConfig(name="math.sqrt"), |
| 217 | + ) |
| 218 | + } |
| 219 | + ) |
| 220 | + |
| 221 | + try: |
| 222 | + result = register_custom_metrics_from_config(eval_config) |
| 223 | + |
| 224 | + assert result is DEFAULT_METRIC_EVALUATOR_REGISTRY |
| 225 | + registered_info = self._registered_metric_info( |
| 226 | + DEFAULT_METRIC_EVALUATOR_REGISTRY, self._CUSTOM_METRIC_NAME |
| 227 | + ) |
| 228 | + assert registered_info.metric_name == self._CUSTOM_METRIC_NAME |
| 229 | + finally: |
| 230 | + DEFAULT_METRIC_EVALUATOR_REGISTRY._registry.pop( |
| 231 | + self._CUSTOM_METRIC_NAME, None |
| 232 | + ) |
| 233 | + |
| 234 | + |
123 | 235 | class TestMetricInfoProviders: |
124 | 236 | """Test cases for MetricInfoProviders.""" |
125 | 237 |
|
|
0 commit comments