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
97105class 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+
150168def _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
727745def _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
842858def _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
891907def _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
949965def _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
10741090def _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
14381456def _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