11import argparse
2+ import asyncio
23import json
34import logging
5+ from typing import Any , Optional
46
5- import requests
7+ # Load environment variables from .env file if available
8+ try :
9+ from dotenv import load_dotenv
610
7- REQUEST_TIMEOUT = 60 * 60 * 5
11+ load_dotenv ()
12+ except ImportError :
13+ logging .warning (
14+ "dotenv not installed, skipping loading environment variables"
15+ )
16+
17+ from google .adk .runners import Runner
18+ from google .adk .sessions .in_memory_session_service import InMemorySessionService
19+ from google .genai .types import Content , Part
820
21+ from auto_agent .agent import root_agent
922
10- # Configure logging
11- logging .basicConfig (
12- level = logging .INFO , format = "%(asctime)s - %(levelname)s - %(message)s"
13- )
1423logger = logging .getLogger (__name__ )
1524
1625
1726class AutoAgentClient :
1827 user_id : str
1928 session_id : str
2029 query : str
21- base_url : str
2230 app_name : str = "auto_agent"
2331
2432 def __init__ (
2533 self ,
2634 user_id : str ,
2735 session_id : str ,
2836 query : str ,
29- base_url : str = "http://localhost:8000" ,
37+ agent : Optional [ Any ] = None ,
3038 ):
3139 self .user_id = user_id
3240 self .session_id = session_id
3341 self .query = query
34- self .base_url = base_url
35-
36- def create_session (self ):
37- session_url = f"{ self .base_url } /apps/{ self .app_name } /users/{ self .user_id } /sessions/{ self .session_id } "
38- response = requests .post (
39- session_url , headers = {"Content-Type" : "application/json" }
40- )
41- return response
42-
43- def send_query (self ):
44- run_url = f"{ self .base_url } /run"
45- payload = {
46- "appName" : self .app_name ,
47- "userId" : self .user_id ,
48- "sessionId" : self .session_id ,
49- "newMessage" : {"role" : "user" , "parts" : [{"text" : self .query }]},
50- }
51- response = requests .post (
52- run_url ,
53- headers = {"Content-Type" : "application/json" },
54- data = json .dumps (payload ),
55- timeout = REQUEST_TIMEOUT ,
42+ self .agent = agent or root_agent
43+ self .session_service = InMemorySessionService ()
44+ self .session = None
45+
46+ async def create_session (
47+ self , initial_state : Optional [dict [str , Any ]] = None
48+ ) -> None :
49+ self .session = await self .session_service .create_session (
50+ app_name = self .app_name ,
51+ user_id = self .user_id ,
52+ session_id = self .session_id ,
53+ state = initial_state ,
5654 )
57- return response
5855
59- def _get_session_data (self ) -> dict :
60- session_url = f"{ self .base_url } /apps/{ self .app_name } /users/{ self .user_id } /sessions/{ self .session_id } "
61- response = requests .get (
62- session_url , headers = {"Content-Type" : "application/json" }
56+ def get_session_data (self ) -> dict :
57+ if not self .session :
58+ raise ValueError ("Session has not been created yet." )
59+
60+ return json .loads (
61+ self .session .model_dump_json (by_alias = True , exclude_none = True )
6362 )
64- if response .status_code != 200 :
65- raise ValueError (f"Failed to get state: { response .status_code } " )
66- try :
67- return response .json ()
68- except json .JSONDecodeError :
69- raise ValueError ("Failed to decode JSON response from server" )
7063
71- def get_state (self , key : str = None ):
72- state_data = self ._get_session_data ()
73- state = state_data .get ("state" )
74- if not state :
75- raise ValueError ("State not found in response" )
64+ def get_state (self , key : Optional [str ] = None ) -> Any :
65+ if not self .session :
66+ raise ValueError ("Session has not been created yet." )
7667
68+ state = self .session .state
7769 if key is None :
7870 return state
7971
8072 if key not in state :
8173 raise ValueError (f"Key '{ key } ' not found in state" )
8274 return state .get (key )
8375
76+ async def run_async (self ) -> None :
77+ if not self .session :
78+ await self .create_session ()
79+
80+ runner = Runner (
81+ app_name = self .app_name ,
82+ agent = self .agent ,
83+ session_service = self .session_service ,
84+ )
8485
85- def run_agent (client : AutoAgentClient ):
86- # Create session
87- session_response = client .create_session ()
88- if session_response .status_code != 200 :
89- raise ValueError (f"Failed to create session: { session_response .text } " )
86+ new_message = Content (parts = [Part (text = self .query )])
9087
91- # Send query
92- query_response = client .send_query ()
93- if query_response .status_code != 200 :
94- raise ValueError (f"ADK server returned error: { query_response .text } " )
88+ logger .info (f"Starting in-process agent run for session { self .session_id } " )
89+ try :
90+ async for event in runner .run_async (
91+ user_id = self .user_id ,
92+ session_id = self .session_id ,
93+ new_message = new_message ,
94+ ):
95+ pass
96+ logger .info (
97+ f"Finished in-process agent run for session { self .session_id } "
98+ )
99+ finally :
100+ # Retrieve the updated session from the session service, even if the run crashed
101+ self .session = await self .session_service .get_session (
102+ app_name = self .app_name ,
103+ user_id = self .user_id ,
104+ session_id = self .session_id ,
105+ )
95106
96107
97108def read_query_from_file (file_path : str ) -> str :
@@ -113,11 +124,6 @@ def main():
113124 default = "client_query.txt" ,
114125 help = "File containing the query to send" ,
115126 )
116- parser .add_argument (
117- "--already-generated" ,
118- action = "store_true" ,
119- help = "Whether to use an already generated script" ,
120- )
121127 args = parser .parse_args ()
122128
123129 user_id = args .user_id
@@ -128,16 +134,10 @@ def main():
128134 # Create client instance
129135 client = AutoAgentClient (user_id , session_id , query )
130136
131- if not args .already_generated :
132- # Generate and return script
133- logger .info (
134- f"Generating script for user { user_id } in session { session_id } with query: { query } "
135- )
136- run_agent (client )
137- else :
138- logger .info (
139- f"Using already generated script for user { user_id } in session { session_id } "
140- )
137+ logger .info (
138+ f"Generating script for user { user_id } in session { session_id } with query: { query } "
139+ )
140+ asyncio .run (client .run_async ())
141141
142142
143143if __name__ == "__main__" :
0 commit comments