Skip to content

Commit 262fbc3

Browse files
vertex-sdk-botcopybara-github
authored andcommitted
feat: Add ADK version check and set MemoryBankService as default when google-adk>=1.5.0
PiperOrigin-RevId: 781290630
1 parent f99b105 commit 262fbc3

2 files changed

Lines changed: 56 additions & 20 deletions

File tree

tests/unit/vertex_adk/test_reasoning_engine_templates_adk.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -84,10 +84,10 @@ def simple_span_processor_mock():
8484

8585

8686
@pytest.fixture
87-
def mock_adk_major_version():
87+
def mock_adk_version():
8888
with mock.patch(
89-
"google.cloud.aiplatform.vertexai.preview.reasoning_engines.templates.adk.get_adk_major_version",
90-
return_value=1,
89+
"google.cloud.aiplatform.vertexai.preview.reasoning_engines.templates.adk.get_adk_version",
90+
return_value="1.5.0",
9191
):
9292
yield
9393

@@ -148,17 +148,17 @@ async def run_async(self, *args, **kwargs):
148148
)
149149

150150

151-
@pytest.mark.usefixtures("google_auth_mock", "mock_adk_major_version")
151+
@pytest.mark.usefixtures("google_auth_mock")
152152
class TestAdkApp:
153-
def test_adk_major_version(self):
153+
def test_adk_version(self):
154154
with mock.patch(
155-
"google.cloud.aiplatform.vertexai.preview.reasoning_engines.templates.adk.get_adk_major_version",
156-
return_value=0,
155+
"google.cloud.aiplatform.vertexai.preview.reasoning_engines.templates.adk.get_adk_version",
156+
return_value="0.5.0",
157157
):
158158
with pytest.raises(
159159
ValueError,
160160
match=(
161-
"Unsupported google-adk major version: 0, please use"
161+
"Unsupported google-adk version: 0.5.0, please use"
162162
" google-adk>=1.0.0 for AdkApp deployment."
163163
),
164164
):
@@ -400,7 +400,7 @@ def test_enable_tracing_warning(self, caplog):
400400
# assert "enable_tracing=True but proceeding with tracing disabled" in caplog.text
401401

402402

403-
@pytest.mark.usefixtures("mock_adk_major_version")
403+
@pytest.mark.usefixtures("mock_adk_version")
404404
class TestAdkAppErrors:
405405
def test_raise_get_session_not_found_error(self):
406406
with pytest.raises(

vertexai/preview/reasoning_engines/templates/adk.py

Lines changed: 47 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -81,14 +81,31 @@
8181
_DEFAULT_USER_ID = "default-user-id"
8282

8383

84-
def get_adk_major_version() -> int:
85-
"""Get the major version of google-adk."""
84+
def get_adk_version() -> Optional[str]:
85+
"""Returns the version of the ADK package."""
8686
try:
8787
from google.adk import version
8888

89-
return int(version.__version__.split(".")[0])
89+
return version.__version__
9090
except ImportError:
91-
return 0
91+
return None
92+
93+
94+
def is_version_sufficient(version_to_check: str) -> bool:
95+
"""Compares the existing version of ADK with the required version.
96+
97+
Args:
98+
version_to_check: The version string to check.
99+
100+
Returns:
101+
True if the existing version is sufficient, False otherwise.
102+
"""
103+
try:
104+
from packaging.version import parse
105+
106+
return parse(get_adk_version()) >= parse(version_to_check)
107+
except (AttributeError, ImportError):
108+
return False
92109

93110

94111
class _ArtifactVersion:
@@ -294,10 +311,10 @@ def __init__(
294311
"""An ADK Application."""
295312
from google.cloud.aiplatform import initializer
296313

297-
adk_major_version = get_adk_major_version()
298-
if adk_major_version < 1:
314+
adk_version = get_adk_version()
315+
if not is_version_sufficient("1.0.0"):
299316
msg = (
300-
f"Unsupported google-adk major version: {adk_major_version}, "
317+
f"Unsupported google-adk version: {adk_version}, "
301318
"please use google-adk>=1.0.0 for AdkApp deployment."
302319
)
303320
raise ValueError(msg)
@@ -464,16 +481,35 @@ def set_up(self):
464481
VertexAiSessionService,
465482
)
466483

467-
self._tmpl_attrs["session_service"] = VertexAiSessionService(
468-
project=project,
469-
location=location,
470-
)
484+
if is_version_sufficient("1.5.0"):
485+
self._tmpl_attrs["session_service"] = VertexAiSessionService(
486+
project=project,
487+
location=location,
488+
agent_engine_id=os.environ.get("GOOGLE_CLOUD_AGENT_ENGINE_ID"),
489+
)
490+
else:
491+
self._tmpl_attrs["session_service"] = VertexAiSessionService(
492+
project=project,
493+
location=location,
494+
)
471495
else:
472496
self._tmpl_attrs["session_service"] = InMemorySessionService()
473497

474498
memory_service_builder = self._tmpl_attrs.get("memory_service_builder")
475499
if memory_service_builder:
476500
self._tmpl_attrs["memory_service"] = memory_service_builder()
501+
elif "GOOGLE_CLOUD_AGENT_ENGINE_ID" in os.environ and is_version_sufficient(
502+
"1.5.0"
503+
):
504+
from google.adk.memory.vertex_ai_memory_bank_service import (
505+
VertexAiMemoryBankService,
506+
)
507+
508+
self._tmpl_attrs["memory_service"] = VertexAiMemoryBankService(
509+
project=project,
510+
location=location,
511+
agent_engine_id=os.environ.get("GOOGLE_CLOUD_AGENT_ENGINE_ID"),
512+
)
477513
else:
478514
self._tmpl_attrs["memory_service"] = InMemoryMemoryService()
479515

0 commit comments

Comments
 (0)