Skip to content

Commit b04a76c

Browse files
authored
Merge pull request #72 from sfw/sfw/surplus-practice-cache
Cache surplus practice blocks for instant delivery
2 parents 675dfb1 + a728cb7 commit b04a76c

4 files changed

Lines changed: 623 additions & 0 deletions

File tree

src/dibble/bootstrap.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
from dibble.services.outcome_store import SQLiteOutcomeStore
3030
from dibble.services.strand_store import SQLiteStrandStore
3131
from dibble.services.generation_engine import GenerationEngine
32+
from dibble.services.surplus_practice_cache import SurplusPracticeCache
3233
from dibble.services.generation_mode_calibration import GenerationModeCalibrator
3334
from dibble.services.generated_content_store import SQLiteGeneratedContentStore
3435
from dibble.services.knowledge_component_store import SQLiteKnowledgeComponentStore
@@ -258,12 +259,17 @@ def build_application_services(
258259
strategy_signal_service=learner_strategy_signal_service,
259260
within_session_adaptation_service=within_session_adaptation_service,
260261
)
262+
surplus_practice_cache = SurplusPracticeCache(
263+
generated_content_store=generated_content_store,
264+
cache_ttl_seconds=settings.generation_cache_ttl_seconds,
265+
)
261266
generation_engine = GenerationEngine(
262267
retriever=plugins.retriever,
263268
router=router_plugin,
264269
provider=plugins.provider,
265270
validator=plugins.validator,
266271
generated_content_store=generated_content_store,
272+
surplus_practice_cache=surplus_practice_cache,
267273
cache_ttl_seconds=settings.generation_cache_ttl_seconds,
268274
)
269275
misconception_remediation_outcome_signal_service = (

src/dibble/services/generation_engine.py

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
from dibble.services.generation_modes import build_generation_mode_plan
3232
from dibble.services.protocols import GeneratedContentStore
3333
from dibble.services.runtime_telemetry import log_runtime_event
34+
from dibble.services.surplus_practice_cache import SurplusPracticeCache
3435

3536
logger = 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

Comments
 (0)