3131from dibble .services .generation_modes import build_generation_mode_plan
3232from dibble .services .protocols import GeneratedContentStore
3333from dibble .services .runtime_telemetry import log_runtime_event
34+ from dibble .services .surplus_practice_cache import SurplusPracticeCache
3435
3536logger = logging .getLogger (__name__ )
3637
@@ -44,6 +45,7 @@ def __init__(
4445 validator : ValidatorPlugin ,
4546 moderation_service : ContentModerationService | None = None ,
4647 generated_content_store : GeneratedContentStore | None = None ,
48+ surplus_practice_cache : SurplusPracticeCache | None = None ,
4749 cache_ttl_seconds : int = 3600 ,
4850 time_provider = monotonic ,
4951 ) -> None :
@@ -53,6 +55,7 @@ def __init__(
5355 self .validator = validator
5456 self .moderation_service = moderation_service or ContentModerationService ()
5557 self .generated_content_store = generated_content_store
58+ self .surplus_practice_cache = surplus_practice_cache
5659 self .cache_ttl_seconds = max (0 , cache_ttl_seconds )
5760 self .time_provider = time_provider
5861
@@ -89,6 +92,10 @@ def generate(
8992 route = route .model_dump (mode = "json" ),
9093 grounding = [item .model_dump (mode = "json" ) for item in grounding ],
9194 )
95+ surplus = self ._pop_surplus (profile , request )
96+ if surplus is not None :
97+ return surplus .response
98+
9299 cache_key = self ._cache_key (profile , request , route , grounding )
93100 cached = self ._get_cached_content (cache_key = cache_key )
94101 if cached is not None :
@@ -104,6 +111,7 @@ def generate(
104111 return cached .response
105112
106113 started_at = self .time_provider ()
114+ surplus_blocks : list [GeneratedBlock ] = []
107115 request_moderation = self .moderation_service .moderate_request (request )
108116 if request_moderation .status == "flagged" :
109117 blocks = self ._moderation_fallback_blocks (
@@ -125,6 +133,7 @@ def generate(
125133 else :
126134 blocks = self .provider .generate (profile , request , route , grounding )
127135 blocks = normalize_generated_blocks (blocks )
136+ blocks , surplus_blocks = self ._split_surplus (blocks )
128137 moderation = self .moderation_service .moderate_blocks (blocks )
129138 if moderation .status == "flagged" :
130139 original_blocks = len (blocks )
@@ -158,6 +167,7 @@ def generate(
158167 ),
159168 )
160169 self ._store_generated_content (cache_key = cache_key , content = content )
170+ self ._cache_surplus (surplus_blocks , blocks , content , profile , request )
161171 log_runtime_event (
162172 logger ,
163173 logging .DEBUG ,
@@ -178,6 +188,31 @@ def stream_generate(
178188 ) -> Iterator [GenerationStreamEvent ]:
179189 grounding = self ._safe_retrieve (profile , request )
180190 route = self .router .route (profile , request )
191+
192+ surplus = self ._pop_surplus (profile , request )
193+ if surplus is not None :
194+ yield GenerationStreamEvent (
195+ event = "start" ,
196+ student_id = profile .student_id ,
197+ route = surplus .response .route ,
198+ grounding = surplus .response .grounding ,
199+ )
200+ for chunk in self ._stream_cached_blocks (surplus .response .blocks ):
201+ yield GenerationStreamEvent (
202+ event = "delta" ,
203+ student_id = profile .student_id ,
204+ chunk = chunk ,
205+ )
206+ yield GenerationStreamEvent (
207+ event = "complete" ,
208+ student_id = profile .student_id ,
209+ route = surplus .response .route ,
210+ grounding = surplus .response .grounding ,
211+ validation_issues = surplus .response .validation_issues ,
212+ response = surplus .response ,
213+ )
214+ return
215+
181216 cache_key = self ._cache_key (profile , request , route , grounding )
182217 cached = self ._get_cached_content (cache_key = cache_key )
183218 if cached is not None :
@@ -213,6 +248,7 @@ def stream_generate(
213248 return
214249
215250 started_at = self .time_provider ()
251+ surplus_blocks : list [GeneratedBlock ] = []
216252 request_moderation = self .moderation_service .moderate_request (request )
217253 if request_moderation .status == "flagged" :
218254 blocks = self ._moderation_fallback_blocks (
@@ -264,6 +300,7 @@ def stream_generate(
264300 blocks = normalize_generated_blocks (
265301 [block_buffers [index ] for index in sorted (block_buffers )]
266302 )
303+ blocks , surplus_blocks = self ._split_surplus (blocks )
267304 moderation = self .moderation_service .moderate_blocks (blocks )
268305 if moderation .status == "flagged" :
269306 original_blocks = len (blocks )
@@ -308,6 +345,7 @@ def stream_generate(
308345 ),
309346 )
310347 self ._store_generated_content (cache_key = cache_key , content = content )
348+ self ._cache_surplus (surplus_blocks , blocks , content , profile , request )
311349 log_runtime_event (
312350 logger ,
313351 logging .DEBUG ,
@@ -330,6 +368,42 @@ def stream_generate(
330368 response = content .response ,
331369 )
332370
371+ def _pop_surplus (
372+ self , profile : LearnerProfile , request : GenerationRequest
373+ ) -> GeneratedContent | None :
374+ if self .surplus_practice_cache is None :
375+ return None
376+ return self .surplus_practice_cache .pop_surplus (
377+ student_id = profile .student_id ,
378+ learning_session_id = request .learning_session_id ,
379+ )
380+
381+ def _split_surplus (
382+ self , blocks : list [GeneratedBlock ]
383+ ) -> tuple [list [GeneratedBlock ], list [GeneratedBlock ]]:
384+ if self .surplus_practice_cache is None :
385+ return blocks , []
386+ return SurplusPracticeCache .split_practice_blocks (blocks )
387+
388+ def _cache_surplus (
389+ self ,
390+ surplus_blocks : list [GeneratedBlock ],
391+ delivery_blocks : list [GeneratedBlock ],
392+ content : GeneratedContent ,
393+ profile : LearnerProfile ,
394+ request : GenerationRequest ,
395+ ) -> None :
396+ if not surplus_blocks or self .surplus_practice_cache is None :
397+ return
398+ non_practice = [b for b in delivery_blocks if b .kind != "practice_problem" ]
399+ self .surplus_practice_cache .cache_surplus (
400+ surplus_blocks = surplus_blocks ,
401+ non_practice_blocks = non_practice ,
402+ parent_content = content ,
403+ profile = profile ,
404+ request = request ,
405+ )
406+
333407 def _build_response (
334408 self ,
335409 profile : LearnerProfile ,
0 commit comments