diff --git a/src/google/adk/evaluation/conversation_scenarios.py b/src/google/adk/evaluation/conversation_scenarios.py index e74ae3b1..dc2d09b7 100644 --- a/src/google/adk/evaluation/conversation_scenarios.py +++ b/src/google/adk/evaluation/conversation_scenarios.py @@ -20,7 +20,7 @@ from pydantic import Field from pydantic import field_validator from .common import EvalBaseModel -from .simulation.pre_built_personas import DEFAULT_USER_PERSONA_REGISTRY +from .simulation.pre_built_personas import get_default_persona_registry from .simulation.user_simulator_personas import UserPersona @@ -62,7 +62,7 @@ class ConversationScenario(EvalBaseModel): cls, value: Optional[UserPersona | str] ) -> Optional[UserPersona]: if value is not None and isinstance(value, str): - return DEFAULT_USER_PERSONA_REGISTRY.get_persona(value) + return get_default_persona_registry().get_persona(value) return value diff --git a/src/google/adk/evaluation/simulation/pre_built_personas.py b/src/google/adk/evaluation/simulation/pre_built_personas.py index 63f3f25e..03f52ae6 100644 --- a/src/google/adk/evaluation/simulation/pre_built_personas.py +++ b/src/google/adk/evaluation/simulation/pre_built_personas.py @@ -509,7 +509,7 @@ class _PreBuiltPersonas(enum.Enum): ) -def _get_default_persona_registry() -> UserPersonaRegistry: +def get_default_persona_registry() -> UserPersonaRegistry: registry = UserPersonaRegistry() registry.register_persona( @@ -523,6 +523,3 @@ def _get_default_persona_registry() -> UserPersonaRegistry: ) return registry - - -DEFAULT_USER_PERSONA_REGISTRY = _get_default_persona_registry() diff --git a/tests/unittests/evaluation/simulation/test_pre_built_personas.py b/tests/unittests/evaluation/simulation/test_pre_built_personas.py index f83fdd88..32401da4 100644 --- a/tests/unittests/evaluation/simulation/test_pre_built_personas.py +++ b/tests/unittests/evaluation/simulation/test_pre_built_personas.py @@ -12,9 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. -from google.adk.evaluation.simulation import pre_built_personas +from google.adk.evaluation.simulation.pre_built_personas import get_default_persona_registry def test_get_default_persona_registry(): """Tests that the default persona registry can be loaded.""" - assert pre_built_personas.DEFAULT_USER_PERSONA_REGISTRY is not None + assert get_default_persona_registry() is not None