Skip to content

Commit 54da0f1

Browse files
committed
Extend preprocessing/postprocessing hooks with Session and Prompt access
- preprocess_request now receives Session for context-aware preprocessing - postprocess_response now receives both Prompt and Session - Preprocessing moved inside prepare() for proper request transformation - Use instance variables to pass context from prepare() to postprocess_response()
1 parent 928b6ee commit 54da0f1

2 files changed

Lines changed: 96 additions & 68 deletions

File tree

src/trivia_agent/worker.py

Lines changed: 39 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -269,29 +269,45 @@ def __init__(
269269
self._workspace_dir = workspace_dir or DEFAULT_WORKSPACE_DIR
270270
self._base_template = build_prompt_template()
271271
self._overrides_store = overrides_store
272+
self._last_prompt: Prompt[TriviaResponse] | None = None
273+
self._last_session: Session | None = None
272274

273-
def preprocess_request(self, request: TriviaRequest) -> TriviaRequest:
275+
def preprocess_request(
276+
self,
277+
request: TriviaRequest,
278+
session: Session,
279+
) -> TriviaRequest:
274280
"""Transform request before agent processing.
275281
276282
Override this method in subclasses to implement custom preprocessing
277-
logic such as validation, normalization, or enrichment.
283+
logic such as validation, normalization, or enrichment. Has access to
284+
the session for context-aware preprocessing.
278285
279286
Args:
280287
request: The incoming TriviaRequest.
288+
session: The session for this execution.
281289
282290
Returns:
283291
TriviaRequest: The preprocessed request.
284292
"""
285293
return request
286294

287-
def postprocess_response(self, response: TriviaResponse) -> TriviaResponse:
295+
def postprocess_response(
296+
self,
297+
response: TriviaResponse,
298+
prompt: Prompt[TriviaResponse],
299+
session: Session,
300+
) -> TriviaResponse:
288301
"""Transform response before returning to caller.
289302
290303
Override this method in subclasses to implement custom postprocessing
291-
logic such as formatting, cleanup, or validation.
304+
logic such as formatting, cleanup, or validation. Has access to the
305+
prompt and session for context-aware postprocessing.
292306
293307
Args:
294308
response: The TriviaResponse from the agent.
309+
prompt: The prompt used for this execution.
310+
session: The session used for this execution.
295311
296312
Returns:
297313
TriviaResponse: The postprocessed response.
@@ -310,38 +326,20 @@ def prepare(
310326
for isolation, builds the complete PromptTemplate with workspace section,
311327
binds request parameters, and optionally applies experiment overrides.
312328
313-
Note: Request preprocessing is applied in the execute() method before
314-
prepare() is called. This ensures the preprocessed request is used
315-
consistently throughout the execution flow.
316-
317-
This method demonstrates key WINK patterns:
318-
319-
- **Session per request**: Each request gets its own Session for proper
320-
isolation between concurrent requests
321-
- **Dynamic workspace section**: ClaudeAgentWorkspaceSection is created
322-
here (not in build_prompt_template) because it needs the Session
323-
- **Parameter binding**: Binds QuestionParams with the user's question
324-
and EmptyParams for parameterless sections
325-
- **Experiment support**: Uses experiment.overrides_tag to select prompt
326-
variants for A/B testing; defaults to "latest"
327-
- **Override seeding**: Automatically seeds the overrides store with
328-
current prompt state, creating editable files for customization
329+
Calls preprocess_request() to allow request transformation before binding.
329330
330331
Args:
331-
request: TriviaRequest containing the question field. The question
332-
is bound to QuestionSection via QuestionParams.
333-
experiment: Optional Experiment instance for evaluation runs. When
334-
provided, uses experiment.overrides_tag to select prompt variant.
335-
Pass None for production requests.
332+
request: TriviaRequest containing the question field.
333+
experiment: Optional Experiment instance for evaluation runs.
336334
337335
Returns:
338-
tuple[Prompt[TriviaResponse], Session]: A 2-tuple containing:
339-
- Prompt: Fully configured prompt with all sections and bound
340-
parameters, ready for adapter.run()
341-
- Session: Fresh session instance for this request's execution
336+
tuple[Prompt[TriviaResponse], Session]: Prompt and session for execution.
342337
"""
343338
session = Session()
344339

340+
# Allow subclasses to preprocess the request
341+
request = self.preprocess_request(request, session)
342+
345343
# Create workspace section with seeded files
346344
# This needs to be per-request because it references the session
347345
workspace_section = create_workspace_section(
@@ -383,6 +381,10 @@ def prepare(
383381
prompt.bind(QuestionParams(question=request.question))
384382
prompt.bind(EmptyParams()) # For sections without params (GameRules, Hints)
385383

384+
# Store for postprocess_response access
385+
self._last_prompt = prompt
386+
self._last_session = session
387+
386388
return prompt, session
387389

388390
def execute(
@@ -397,8 +399,8 @@ def execute(
397399
) -> tuple[PromptResponse[TriviaResponse], Session]:
398400
"""Execute a trivia request with preprocessing and postprocessing.
399401
400-
Overrides the parent MainLoop.execute() to apply preprocess_request()
401-
before execution and postprocess_response() after execution.
402+
Preprocessing happens in prepare() via preprocess_request().
403+
Postprocessing happens after execution via postprocess_response().
402404
403405
Args:
404406
request: TriviaRequest containing the question to process.
@@ -411,12 +413,9 @@ def execute(
411413
Returns:
412414
tuple[PromptResponse[TriviaResponse], Session]: Response and session.
413415
"""
414-
# Apply preprocessing to the request
415-
preprocessed_request = self.preprocess_request(request)
416-
417-
# Execute with the preprocessed request
416+
# Execute (prepare() is called internally by parent, which calls preprocess_request)
418417
prompt_response, session = super().execute(
419-
preprocessed_request,
418+
request,
420419
budget=budget,
421420
deadline=deadline,
422421
resources=resources,
@@ -426,8 +425,10 @@ def execute(
426425

427426
# Apply postprocessing to the response output (if present)
428427
output = prompt_response.output
429-
if output is not None:
430-
postprocessed_output = self.postprocess_response(output)
428+
if output is not None and self._last_prompt is not None:
429+
postprocessed_output = self.postprocess_response(
430+
output, self._last_prompt, self._last_session or session
431+
)
431432
return prompt_response.update(output=postprocessed_output), session # type: ignore[return-value]
432433

433434
return prompt_response, session

tests/trivia_agent/test_worker.py

Lines changed: 57 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -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

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

Comments
 (0)