Skip to content

Commit 6b177ca

Browse files
yeesiancopybara-github
authored andcommitted
fix: Add support for ADK Agents as a supported type for agent engine
PiperOrigin-RevId: 769622242
1 parent b57cbd3 commit 6b177ca

1 file changed

Lines changed: 31 additions & 13 deletions

File tree

vertexai/agent_engines/_agent_engines.py

Lines changed: 31 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,14 @@
9393
_DEFAULT_AGENT_FRAMEWORK = "custom"
9494

9595

96+
try:
97+
from google.adk.agents import BaseAgent
98+
99+
ADKAgent = BaseAgent
100+
except (ImportError, AttributeError):
101+
ADKAgent = None
102+
103+
96104
@typing.runtime_checkable
97105
class Queryable(Protocol):
98106
"""Protocol for Agent Engines that can be queried."""
@@ -147,6 +155,16 @@ def register_operations(self, **kwargs):
147155
"""Register the user provided operations (modes and methods)."""
148156

149157

158+
_AgentEngineInterface = Union[
159+
ADKAgent,
160+
AsyncQueryable,
161+
AsyncStreamQueryable,
162+
OperationRegistrable,
163+
Queryable,
164+
StreamQueryable,
165+
]
166+
167+
150168
def _wrap_agent_operation(agent: Any, operation: str):
151169
"""Wraps an agent operation into a method (works for all API modes)."""
152170

@@ -294,7 +312,7 @@ def resource_name(self) -> str:
294312
@classmethod
295313
def create(
296314
cls,
297-
agent_engine: Optional[Union[Queryable, OperationRegistrable]] = None,
315+
agent_engine: Optional[_AgentEngineInterface] = None,
298316
*,
299317
requirements: Optional[Union[str, Sequence[str]]] = None,
300318
display_name: Optional[str] = None,
@@ -503,7 +521,7 @@ def create(
503521
def update(
504522
self,
505523
*,
506-
agent_engine: Optional[Union[Queryable, OperationRegistrable]] = None,
524+
agent_engine: Optional[_AgentEngineInterface] = None,
507525
requirements: Optional[Union[str, Sequence[str]]] = None,
508526
display_name: Optional[str] = None,
509527
description: Optional[str] = None,
@@ -725,10 +743,8 @@ def _validate_staging_bucket_or_raise(staging_bucket: str) -> str:
725743

726744

727745
def _validate_agent_engine_or_raise(
728-
agent_engine: Union[
729-
Queryable, OperationRegistrable, StreamQueryable, AsyncStreamQueryable
730-
],
731-
) -> Union[Queryable, OperationRegistrable, StreamQueryable, AsyncStreamQueryable]:
746+
agent_engine: _AgentEngineInterface,
747+
) -> _AgentEngineInterface:
732748
"""Tries to validate the agent engine.
733749
734750
The agent engine must have one of the following:
@@ -841,7 +857,7 @@ def _validate_agent_engine_or_raise(
841857

842858
def _validate_requirements_or_raise(
843859
*,
844-
agent_engine: Union[Queryable, OperationRegistrable],
860+
agent_engine: _AgentEngineInterface,
845861
requirements: Optional[Sequence[str]] = None,
846862
) -> Sequence[str]:
847863
"""Tries to validate the requirements."""
@@ -890,7 +906,7 @@ def _get_gcs_bucket(
890906

891907
def _upload_agent_engine(
892908
*,
893-
agent_engine: Union[Queryable, OperationRegistrable],
909+
agent_engine: _AgentEngineInterface,
894910
gcs_bucket: storage.Bucket,
895911
gcs_dir_name: str,
896912
) -> None:
@@ -947,7 +963,7 @@ def _upload_extra_packages(
947963

948964

949965
def _prepare(
950-
agent_engine: Optional[Union[Queryable, OperationRegistrable]],
966+
agent_engine: Optional[_AgentEngineInterface],
951967
requirements: Optional[Sequence[str]],
952968
extra_packages: Optional[Sequence[str]],
953969
project: str,
@@ -1072,7 +1088,7 @@ def _generate_deployment_spec_or_raise(
10721088

10731089

10741090
def _get_agent_framework(
1075-
agent_engine: Union[Queryable, OperationRegistrable],
1091+
agent_engine: _AgentEngineInterface,
10761092
) -> str:
10771093
if (
10781094
hasattr(agent_engine, _AGENT_FRAMEWORK_ATTR)
@@ -1087,7 +1103,7 @@ def _generate_update_request_or_raise(
10871103
resource_name: str,
10881104
staging_bucket: str,
10891105
gcs_dir_name: str = _DEFAULT_GCS_DIR_NAME,
1090-
agent_engine: Optional[Union[Queryable, OperationRegistrable]] = None,
1106+
agent_engine: Optional[_AgentEngineInterface] = None,
10911107
requirements: Optional[Union[str, Sequence[str]]] = None,
10921108
extra_packages: Optional[Sequence[str]] = None,
10931109
display_name: Optional[str] = None,
@@ -1418,7 +1434,9 @@ def _register_api_methods_or_raise(obj: "AgentEngine"):
14181434
setattr(obj, method_name, types.MethodType(method, obj))
14191435

14201436

1421-
def _get_registered_operations(agent_engine: Any) -> Dict[str, List[str]]:
1437+
def _get_registered_operations(
1438+
agent_engine: _AgentEngineInterface,
1439+
) -> Dict[str, List[str]]:
14221440
"""Retrieves registered operations for a AgentEngine."""
14231441
if isinstance(agent_engine, OperationRegistrable):
14241442
return agent_engine.register_operations()
@@ -1436,7 +1454,7 @@ def _get_registered_operations(agent_engine: Any) -> Dict[str, List[str]]:
14361454

14371455

14381456
def _generate_class_methods_spec_or_raise(
1439-
*, agent_engine: Any, operations: Dict[str, List[str]]
1457+
*, agent_engine: _AgentEngineInterface, operations: Dict[str, List[str]]
14401458
) -> List[proto.Message]:
14411459
"""Generates a ReasoningEngineSpec based on the registered operations.
14421460

0 commit comments

Comments
 (0)