77import threading
88from pathlib import Path
99
10- try :
11- import httpx
12- from dotenv import load_dotenv
13- from openai import OpenAI
14- except Exception as e :
15- print ("WHATTTTT : " , e )
16- sys .exit (0 )
10+ import httpx
11+ from dotenv import load_dotenv
12+ from openai import OpenAI
1713
18- try :
19- ROOT = Path (__file__ ).parent .resolve ()
20- if str (ROOT ) not in sys .path :
21- sys .path .insert (0 , str (ROOT ))
14+ ROOT = Path (__file__ ).parent .resolve ()
15+ if str (ROOT ) not in sys .path :
16+ sys .path .insert (0 , str (ROOT ))
2217
23- except Exception as e :
24- print ("BRUH : " , e )
25- sys .exit (0 )
18+ TASK_DIR = ROOT / "tasks"
2619
20+ from cloudenv .client import ClientEnvironment # noqa
21+ from cloudenv .models import CloudEnvAction , CloudEnvObservation # noqa
2722
28- try :
29- TASK_DIR = ROOT / "tasks"
30- from cloudenv .client import ClientEnvironment
31- from cloudenv .models import CloudEnvAction , CloudEnvObservation
23+ load_dotenv ()
3224
33- load_dotenv ()
25+ HF_TOKEN = os .getenv ("HF_TOKEN" ) or os .getenv ("API_KEY" )
26+ API_BASE_URL = os .getenv ("API_BASE_URL" , "https://api.openai.com/v1" )
27+ MODEL_NAME = os .getenv ("MODEL_NAME" , "gpt-4o" )
3428
35- HF_TOKEN = os .getenv ("HF_TOKEN" ) or os . getenv ( "API_KEY" )
36- API_BASE_URL = os .getenv ("API_BASE_URL " , "https://api.openai.com/v1" )
37- MODEL_NAME = os .getenv ("MODEL_NAME " , "gpt-4o" )
29+ MAX_STEPS = int ( os .getenv ("MAX_STEPS" , "20" ) )
30+ TEMPERATURE = float ( os .getenv ("TEMPERATURE " , "0.3" ) )
31+ MAX_TOKENS = int ( os .getenv ("MAX_TOKENS " , "300" ) )
3832
39- MAX_STEPS = int (os .getenv ("MAX_STEPS" , "20" ))
40- TEMPERATURE = float (os .getenv ("TEMPERATURE" , "0.3" ))
41- MAX_TOKENS = int (os .getenv ("MAX_TOKENS" , "300" ))
33+ SERVER_URL = "http://127.0.0.1:8000"
4234
43- SERVER_URL = "http://127.0.0.1:8000"
35+ TASK_ORDER = [ "easy-task" , "medium-task" , "hard-task" ]
4436
45- TASK_ORDER = [ "easy-task" , "medium-task" , "hard-task" ]
37+ client = OpenAI ( base_url = API_BASE_URL , api_key = HF_TOKEN )
4638
47- client = OpenAI (base_url = API_BASE_URL , api_key = HF_TOKEN )
48- SYSTEM_PROMPT = 'Return ONLY JSON: {"service":"","operation":"","instance_id":null,"payload":{}}'
49- except Exception as e :
50- print ("At this point, I give up : " , e )
51- sys .exit (0 )
39+ SYSTEM_PROMPT = '{"service":"","operation":"","instance_id":null,"payload":{}}'
5240
5341
5442def start_ministack ():
55- try :
56- proc = subprocess .Popen (
57- ["uv" , "run" , "ministack" ],
58- cwd = ROOT ,
59- stdout = subprocess .DEVNULL ,
60- stderr = subprocess .DEVNULL ,
61- )
62- time .sleep (3 )
63- return proc
64- except Exception as e :
65- print (f"[MINISTACK START ERROR] { e } " , flush = True )
66- raise
43+ proc = subprocess .Popen (
44+ ["uv" , "run" , "ministack" ],
45+ cwd = ROOT ,
46+ stdout = subprocess .DEVNULL ,
47+ stderr = subprocess .DEVNULL ,
48+ )
49+ time .sleep (3 )
50+ return proc
6751
6852
6953def start_server ():
70- try :
71- proc = subprocess .Popen (
72- ["uv" , "run" , "server" ],
73- cwd = ROOT ,
74- stdout = subprocess .PIPE ,
75- stderr = subprocess .STDOUT ,
76- text = True ,
77- bufsize = 1 ,
78- )
54+ proc = subprocess .Popen (
55+ ["uv" , "run" , "server" ],
56+ cwd = ROOT ,
57+ stdout = subprocess .PIPE ,
58+ stderr = subprocess .STDOUT ,
59+ text = True ,
60+ bufsize = 1 ,
61+ )
7962
80- def stream ():
81- try :
82- for line in proc .stdout :
83- print (f"[SERVER] { line } " , end = "" )
84- except Exception as e :
85- print (f"[SERVER STREAM ERROR] { e } " , flush = True )
63+ def stream ():
64+ for line in proc .stdout :
65+ print (f"[SERVER] { line } " , end = "" )
8666
87- threading .Thread (target = stream , daemon = True ).start ()
67+ threading .Thread (target = stream , daemon = True ).start ()
8868
89- for _ in range (40 ):
90- try :
91- r = httpx .get (f"{ SERVER_URL } /docs" , timeout = 1 )
92- if r .status_code == 200 :
93- print ("[INFO] server ready" , flush = True )
94- return proc
95- except Exception :
96- time .sleep (1 )
69+ for _ in range (40 ):
70+ r = httpx .get (f"{ SERVER_URL } /docs" , timeout = 1 )
71+ if r .status_code == 200 :
72+ print ("[INFO] server ready" )
73+ return proc
74+ time .sleep (1 )
9775
98- raise RuntimeError ("server failed to start" )
99-
100- except Exception as e :
101- print (f"[SERVER START ERROR] { e } " , flush = True )
102- raise
76+ raise RuntimeError ("server failed" )
10377
10478
10579def load_task (name : str ):
106- try :
107- return json .loads ((TASK_DIR / f"{ name } .json" ).read_text ())
108- except Exception as e :
109- print (f"[TASK LOAD ERROR] { name } : { e } " , flush = True )
110- raise
80+ return json .loads ((TASK_DIR / f"{ name } .json" ).read_text ())
11181
11282
11383def register_pipeline (task ):
114- try :
115- url = f"{ SERVER_URL } /api/pipelines"
116- with httpx .Client (timeout = 20 ) as http :
117- r = http .post (url , json = task )
84+ url = f"{ SERVER_URL } /api/pipelines"
85+ with httpx .Client (timeout = 20 ) as http :
86+ r = http .post (url , json = task )
11887
119- if r .status_code == 400 :
120- try :
121- detail = r .json ().get ("detail" , "" )
122- if "already exists" in detail :
123- http .delete (f"{ url } /{ task ['pipeline_id' ]} " )
124- http .post (url , json = task )
125- except Exception :
126- pass
88+ if r .status_code == 400 :
89+ detail = r .json ().get ("detail" , "" )
90+ if "already exists" in detail :
91+ http .delete (f"{ url } /{ task ['pipeline_id' ]} " )
92+ http .post (url , json = task )
12793
128- r .raise_for_status ()
129- except Exception as e :
130- print (f"[REGISTER ERROR] { e } " , flush = True )
94+ r .raise_for_status ()
13195
13296
13397def parse_action (raw : str , task_id : str , episode_id : str ) -> CloudEnvAction :
134- try :
135- raw = raw .strip ()
136- raw = raw [raw .find ("{" ) : raw .rfind ("}" ) + 1 ]
137- d = json .loads (raw )
138- return CloudEnvAction (
139- task_id = task_id ,
140- episode_id = episode_id ,
141- service = d .get ("service" , "internal" ),
142- instance_id = d .get ("instance_id" ),
143- operation = d .get ("operation" , "GetCallerIdentity" ),
144- payload = d .get ("payload" , {}),
145- principal_arn = d .get ("principal_arn" ),
146- intent = d .get ("intent" ),
147- )
148- except Exception :
149- return CloudEnvAction (
150- task_id = task_id ,
151- episode_id = episode_id ,
152- service = "internal" ,
153- operation = "GetCallerIdentity" ,
154- payload = {},
155- )
98+ raw = raw .strip ()
99+ raw = raw [raw .find ("{" ) : raw .rfind ("}" ) + 1 ]
100+ d = json .loads (raw )
101+
102+ return CloudEnvAction (
103+ task_id = task_id ,
104+ episode_id = episode_id ,
105+ service = d .get ("service" , "internal" ),
106+ instance_id = d .get ("instance_id" ),
107+ operation = d .get ("operation" , "GetCallerIdentity" ),
108+ payload = d .get ("payload" , {}),
109+ principal_arn = d .get ("principal_arn" ),
110+ intent = d .get ("intent" ),
111+ )
156112
157113
158114def call_model (step , state , broken , ideal , history ):
159- try :
160- prompt = f"""
115+ prompt = f"""
161116STEP { step }
162117STATE { json .dumps (state )[:1000 ]}
163118BROKEN { json .dumps (broken )[:1000 ]}
164119IDEAL { json .dumps (ideal )[:1000 ]}
165120HISTORY { history [- 5 :]}
166121"""
167- res = client .chat .completions .create (
168- model = MODEL_NAME ,
169- messages = [
170- {"role" : "system" , "content" : SYSTEM_PROMPT },
171- {"role" : "user" , "content" : prompt },
172- ],
173- temperature = TEMPERATURE ,
174- max_tokens = MAX_TOKENS ,
175- )
176- return (res .choices [0 ].message .content or "" ).strip ()
177- except Exception as e :
178- print (f"[MODEL ERROR] { e } " , flush = True )
179- return "{}"
180-
181-
182- async def run_episode (task ):
183- try :
184- pipeline_id = task ["pipeline_id" ]
185- task_id = pipeline_id
186- episode_id = f"{ pipeline_id } -episode-1"
187-
188- env = ClientEnvironment (base_url = SERVER_URL )
189-
190- try :
191- register_pipeline (task )
192- except Exception as e :
193- print (f"[REGISTER ERROR] { e } " , flush = True )
194-
195- try :
196- result = await env .initialize_environment (
197- task_id = task_id ,
198- episode_id = episode_id ,
199- )
200- except Exception as e :
201- print (f"[INIT ERROR] { e } " , flush = True )
202- return
203122
204- obs = result .observation
123+ res = client .chat .completions .create (
124+ model = MODEL_NAME ,
125+ messages = [
126+ {"role" : "system" , "content" : SYSTEM_PROMPT },
127+ {"role" : "user" , "content" : prompt },
128+ ],
129+ temperature = TEMPERATURE ,
130+ max_tokens = MAX_TOKENS ,
131+ )
205132
206- history = []
207- rewards = []
133+ return (res .choices [0 ].message .content or "" ).strip ()
208134
209- for step in range (MAX_STEPS ):
210- try :
211- if obs .done :
212- break
213135
214- raw = call_model (
215- step ,
216- obs .current_pipeline_state ,
217- task ["broken_pipeline" ],
218- task ["ideal_pipeline" ],
219- history ,
220- )
136+ def run_episode (task ):
137+ pipeline_id = task ["pipeline_id" ]
138+ task_id = pipeline_id
139+ episode_id = f"{ pipeline_id } -episode-1"
221140
222- action = parse_action ( raw , task_id , episode_id )
141+ env = ClientEnvironment ( base_url = SERVER_URL )
223142
224- result = await env .step (action )
225- obs = result .observation
143+ register_pipeline (task )
226144
227- rewards .append (float (obs .reward or 0 ))
228- history .append (raw )
145+ result = asyncio .run (
146+ env .initialize_environment (task_id = task_id , episode_id = episode_id )
147+ )
229148
230- except Exception as e :
231- print (f"[STEP ERROR] { e } " , flush = True )
232- break
149+ obs = result .observation
150+
151+ history = []
152+ rewards = []
153+
154+ for step in range (MAX_STEPS ):
155+ if obs .done :
156+ break
157+
158+ raw = call_model (
159+ step ,
160+ obs .current_pipeline_state ,
161+ task ["broken_pipeline" ],
162+ task ["ideal_pipeline" ],
163+ history ,
164+ )
165+
166+ action = parse_action (raw , task_id , episode_id )
167+
168+ result = asyncio .run (env .step (action ))
169+ obs = result .observation
233170
234- score = sum ( rewards ) / ( len ( rewards ) or 1 )
235- print ( f"[END] task= { task_id } score= { score :.3f } " , flush = True )
171+ rewards . append ( float ( obs . reward or 0 ) )
172+ history . append ( raw )
236173
237- except Exception as e :
238- print (f"[EPISODE ERROR] { e } " , flush = True )
174+ score = sum ( rewards ) / ( len ( rewards ) or 1 )
175+ print (f"[END] task= { task_id } score= { score :.3f } " )
239176
240177
241- async def main ():
242- try :
243- miniproc = start_ministack ()
244- serverproc = start_server ()
178+ def main ():
179+ miniproc = start_ministack ()
180+ serverproc = start_server ()
245181
246- try :
247- for name in TASK_ORDER :
248- try :
249- task = load_task (name )
250- await run_episode (task )
251- except Exception as e :
252- print (f"[TASK ERROR] { name } : { e } " , flush = True )
253- finally :
254- try :
255- miniproc .terminate ()
256- except Exception :
257- pass
258- try :
259- serverproc .terminate ()
260- except Exception :
261- pass
182+ for name in TASK_ORDER :
183+ task = load_task (name )
184+ run_episode (task )
262185
263- except Exception as e :
264- print ( f"[FATAL ERROR] { e } " , flush = True )
186+ miniproc . terminate ()
187+ serverproc . terminate ( )
265188
266189
267190if __name__ == "__main__" :
268- try :
269- asyncio .run (main ())
270- except Exception as e :
271- print (f"[UNHANDLED TOP LEVEL ERROR] { e } " , flush = True )
191+ main ()
0 commit comments