Skip to content

Commit 4dbd692

Browse files
Merge pull request #63 from AI-Hypercomputer/shangkun-auto-search
Implement the graph, worker and orchestrator base classes and parallel search and beam search algorithms
2 parents 3f62918 + 0cc34bd commit 4dbd692

12 files changed

Lines changed: 1866 additions & 179 deletions

File tree

MaxKernel/auto_agent/agent_client/auto_agent_client.py

Lines changed: 70 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -1,97 +1,108 @@
11
import argparse
2+
import asyncio
23
import json
34
import 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-
)
1423
logger = logging.getLogger(__name__)
1524

1625

1726
class 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

97108
def 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

143143
if __name__ == "__main__":

0 commit comments

Comments
 (0)