@@ -246,6 +246,8 @@ def test_preprocess_request_returns_unchanged_by_default(
246246 fake_mailboxes : TriviaMailboxes ,
247247 ) -> None :
248248 """Test that preprocess_request returns request unchanged by default."""
249+ from weakincentives .runtime import Session
250+
249251 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
250252
251253 loop = TriviaAgentLoop (
@@ -254,7 +256,8 @@ def test_preprocess_request_returns_unchanged_by_default(
254256 )
255257
256258 request = TriviaRequest (question = "test question" )
257- result = loop .preprocess_request (request )
259+ session = Session ()
260+ result = loop .preprocess_request (request , session )
258261
259262 assert result is request
260263
@@ -263,9 +266,10 @@ def test_preprocess_request_can_be_overridden(
263266 fake_mailboxes : TriviaMailboxes ,
264267 ) -> None :
265268 """Test that preprocess_request can be overridden in subclass."""
269+ from weakincentives .runtime import Session
266270
267271 class CustomLoop (TriviaAgentLoop ):
268- def preprocess_request (self , request : TriviaRequest ) -> TriviaRequest :
272+ def preprocess_request (self , request : TriviaRequest , session : Session ) -> TriviaRequest :
269273 return TriviaRequest (question = request .question .upper ())
270274
271275 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
@@ -276,7 +280,8 @@ def preprocess_request(self, request: TriviaRequest) -> TriviaRequest:
276280 )
277281
278282 request = TriviaRequest (question = "hello world" )
279- result = loop .preprocess_request (request )
283+ session = Session ()
284+ result = loop .preprocess_request (request , session )
280285
281286 assert result .question == "HELLO WORLD"
282287
@@ -289,6 +294,8 @@ def test_postprocess_response_returns_unchanged_by_default(
289294 fake_mailboxes : TriviaMailboxes ,
290295 ) -> None :
291296 """Test that postprocess_response returns response unchanged by default."""
297+ from weakincentives .runtime import Session
298+
292299 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
293300
294301 loop = TriviaAgentLoop (
@@ -297,7 +304,9 @@ def test_postprocess_response_returns_unchanged_by_default(
297304 )
298305
299306 response = TriviaResponse (answer = "42" )
300- result = loop .postprocess_response (response )
307+ mock_prompt = MagicMock ()
308+ session = Session ()
309+ result = loop .postprocess_response (response , mock_prompt , session )
301310
302311 assert result is response
303312
@@ -306,9 +315,16 @@ def test_postprocess_response_can_be_overridden(
306315 fake_mailboxes : TriviaMailboxes ,
307316 ) -> None :
308317 """Test that postprocess_response can be overridden in subclass."""
318+ from weakincentives import Prompt
319+ from weakincentives .runtime import Session
309320
310321 class CustomLoop (TriviaAgentLoop ):
311- def postprocess_response (self , response : TriviaResponse ) -> TriviaResponse :
322+ def postprocess_response (
323+ self ,
324+ response : TriviaResponse ,
325+ prompt : Prompt [TriviaResponse ],
326+ session : Session ,
327+ ) -> TriviaResponse :
312328 return TriviaResponse (answer = f"Answer: { response .answer } " )
313329
314330 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
@@ -319,23 +335,28 @@ def postprocess_response(self, response: TriviaResponse) -> TriviaResponse:
319335 )
320336
321337 response = TriviaResponse (answer = "42" )
322- result = loop .postprocess_response (response )
338+ mock_prompt = MagicMock ()
339+ session = Session ()
340+ result = loop .postprocess_response (response , mock_prompt , session )
323341
324342 assert result .answer == "Answer: 42"
325343
326344
327345class TestTriviaAgentLoopExecute :
328346 """Tests for TriviaAgentLoop.execute() with preprocessing/postprocessing."""
329347
330- def test_execute_calls_preprocess_request (
348+ def test_execute_calls_preprocess_request_via_prepare (
331349 self ,
332350 fake_mailboxes : TriviaMailboxes ,
333351 ) -> None :
334- """Test that execute() calls preprocess_request."""
335- from weakincentives .runtime import MainLoop
352+ """Test that prepare() calls preprocess_request during execution."""
353+ from weakincentives .runtime import Session
354+
355+ preprocess_calls : list [TriviaRequest ] = []
336356
337357 class CustomLoop (TriviaAgentLoop ):
338- def preprocess_request (self , request : TriviaRequest ) -> TriviaRequest :
358+ def preprocess_request (self , request : TriviaRequest , session : Session ) -> TriviaRequest :
359+ preprocess_calls .append (request )
339360 return TriviaRequest (question = request .question .strip ())
340361
341362 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
@@ -345,32 +366,33 @@ def preprocess_request(self, request: TriviaRequest) -> TriviaRequest:
345366 requests = fake_mailboxes .requests ,
346367 )
347368
348- mock_response = MagicMock ()
349- mock_response .output = TriviaResponse (answer = "42" )
350- mock_session = MagicMock ()
351-
352- captured_requests : list [TriviaRequest ] = []
353-
354- def capture_execute (self_arg , request , ** kwargs ):
355- captured_requests .append (request )
356- return (mock_response , mock_session )
357-
358- with patch .object (MainLoop , "execute" , capture_execute ):
359- request = TriviaRequest (question = " What is the answer? " )
360- loop .execute (request )
369+ request = TriviaRequest (question = " What is the answer? " )
370+ prompt , session = loop .prepare (request )
361371
362- assert len (captured_requests ) == 1
363- assert captured_requests [0 ].question == "What is the answer?"
372+ assert len (preprocess_calls ) == 1
373+ assert preprocess_calls [0 ].question == " What is the answer? "
374+ # The prompt should have the preprocessed question bound
375+ rendered = str (prompt .render ())
376+ assert "What is the answer?" in rendered
364377
365378 def test_execute_calls_postprocess_response (
366379 self ,
367380 fake_mailboxes : TriviaMailboxes ,
368381 ) -> None :
369382 """Test that execute() calls postprocess_response."""
370- from weakincentives .runtime import MainLoop
383+ from weakincentives import Prompt
384+ from weakincentives .runtime import MainLoop , Session
385+
386+ postprocess_calls : list [tuple [TriviaResponse , bool , bool ]] = []
371387
372388 class CustomLoop (TriviaAgentLoop ):
373- def postprocess_response (self , response : TriviaResponse ) -> TriviaResponse :
389+ def postprocess_response (
390+ self ,
391+ response : TriviaResponse ,
392+ prompt : Prompt [TriviaResponse ],
393+ session : Session ,
394+ ) -> TriviaResponse :
395+ postprocess_calls .append ((response , prompt is not None , session is not None ))
374396 return TriviaResponse (answer = response .answer .strip ())
375397
376398 mock_adapter : ProviderAdapter [TriviaResponse ] = MagicMock ()
@@ -385,13 +407,18 @@ def postprocess_response(self, response: TriviaResponse) -> TriviaResponse:
385407 mock_response .update = MagicMock (return_value = mock_response )
386408 mock_session = MagicMock ()
387409
410+ # Call prepare() first to set up _last_prompt/_last_session
411+ request = TriviaRequest (question = "What is the answer?" )
412+ loop .prepare (request )
413+
388414 with patch .object (MainLoop , "execute" , return_value = (mock_response , mock_session )):
389- request = TriviaRequest (question = "What is the answer?" )
390415 loop .execute (request )
391416
417+ assert len (postprocess_calls ) == 1
418+ assert postprocess_calls [0 ][0 ].answer == " 42 "
419+ assert postprocess_calls [0 ][1 ] is True # prompt was passed
420+ assert postprocess_calls [0 ][2 ] is True # session was passed
392421 mock_response .update .assert_called_once ()
393- call_kwargs = mock_response .update .call_args .kwargs
394- assert call_kwargs ["output" ].answer == "42"
395422
396423 def test_execute_skips_postprocessing_when_output_is_none (
397424 self ,
0 commit comments