1818from events import EventBus
1919from models .base import MemoryRecord , normalize_modality
2020from models .episodic import EpisodicMemory
21+ from models .procedural import ProceduralMemory
2122from models .semantic import SemanticMemory
2223from retrieval .retriever import UnifiedRetriever
2324from stores .episodic_store import EpisodicStore , EpisodicStoreError , MediaTooLargeError
2425from stores .media_store import MediaStore
26+ from stores .procedural_store import ProceduralMatch , ProceduralStore
2527from stores .semantic_store import SemanticStore
2628from utils .embeddings import EmbeddingProviderError , GeminiEmbedder , TextEmbedder
2729
@@ -141,14 +143,33 @@ def _validate_related_ids(raw_related_ids: Any) -> list[str]:
141143 return related_ids
142144
143145
146+ def _validate_string_list (raw_value : Any , * , field_name : str , required : bool ) -> list [str ]:
147+ if raw_value is None :
148+ values : list [Any ] = []
149+ elif not isinstance (raw_value , list ):
150+ raise ValueError (f"{ field_name } must be an array of strings" )
151+ else :
152+ values = raw_value
153+
154+ if required and not values :
155+ raise ValueError (f"{ field_name } must contain at least one entry" )
156+
157+ cleaned : list [str ] = []
158+ for value in values :
159+ if not isinstance (value , str ) or not value .strip ():
160+ raise ValueError (f"{ field_name } must contain only non-blank strings" )
161+ cleaned .append (value )
162+ return cleaned
163+
164+
144165def _parse_memory_types (raw_value : str | None ) -> list [str ] | None :
145166 if raw_value is None :
146167 return None
147168 values = [value .strip () for value in raw_value .split ("," ) if value .strip ()]
148169 if not values :
149170 return None
150171
151- supported = {"semantic" , "episodic" }
172+ supported = {"semantic" , "episodic" , "procedural" }
152173 invalid = [value for value in values if value not in supported ]
153174 if invalid :
154175 raise ValueError (
@@ -215,6 +236,18 @@ def _serialise_record(record: MemoryRecord) -> dict[str, Any]:
215236 "source_mime_type" : record .source_mime_type ,
216237 }
217238 )
239+ if isinstance (record , ProceduralMemory ):
240+ payload .update (
241+ {
242+ "steps" : record .steps ,
243+ "preconditions" : record .preconditions ,
244+ "success_count" : record .success_count ,
245+ "failure_count" : record .failure_count ,
246+ "total_outcomes" : record .total_outcomes ,
247+ "success_rate" : record .success_rate ,
248+ "wilson_score" : record .wilson_score ,
249+ }
250+ )
218251 return payload
219252
220253
@@ -228,6 +261,15 @@ def _serialise_ranked_result(result) -> dict[str, Any]:
228261 }
229262
230263
264+ def _serialise_procedural_match (match : ProceduralMatch ) -> dict [str , Any ]:
265+ return {
266+ "record" : _serialise_record (match .record ),
267+ "similarity" : match .similarity ,
268+ "wilson_score" : match .wilson_score ,
269+ "combined_score" : match .combined_score ,
270+ }
271+
272+
231273class EventRecorder :
232274 def __init__ (self , bus : EventBus , * , max_events : int = 200 ):
233275 self ._events : deque [dict [str , Any ]] = deque (maxlen = max_events )
@@ -272,10 +314,19 @@ def __init__(
272314 embedder = self .embedder ,
273315 media_store = self .media_store ,
274316 )
317+ self .procedural_store = ProceduralStore (
318+ event_bus = self .bus ,
319+ embedder = self .embedder ,
320+ media_store = self .media_store ,
321+ )
275322 finally :
276323 config .CHROMA_DB_PATH = original_chroma_path
277324 self .retriever = UnifiedRetriever (
278- stores = {"semantic" : self .semantic_store , "episodic" : self .episodic_store },
325+ stores = {
326+ "semantic" : self .semantic_store ,
327+ "episodic" : self .episodic_store ,
328+ "procedural" : self .procedural_store ,
329+ },
279330 event_bus = self .bus ,
280331 )
281332 self .events = EventRecorder (self .bus )
@@ -302,10 +353,12 @@ def save_query_upload(self, upload: UploadFile) -> tuple[str, str]:
302353 def overview (self ) -> dict [str , Any ]:
303354 semantic_count = self .semantic_store ._collection .count ()
304355 episodic_count = self .episodic_store ._collection .count ()
356+ procedural_count = self .procedural_store ._collection .count ()
305357 recent = self .episodic_store .get_recent (5 )
306358 return {
307359 "semantic_count" : semantic_count ,
308360 "episodic_count" : episodic_count ,
361+ "procedural_count" : procedural_count ,
309362 "recent_sessions" : sorted ({record .session_id for record in recent }),
310363 "latest_events" : self .events .snapshot (10 ),
311364 }
@@ -489,6 +542,125 @@ async def create_file_episode(
489542 raise HTTPException (status_code = 422 , detail = str (exc )) from exc
490543 return {"record" : _serialise_record (record )}
491544
545+ @app .post ("/api/memories/procedural" )
546+ async def create_procedural_memory (payload : dict [str , Any ]) -> dict [str , Any ]:
547+ if not payload .get ("content" ):
548+ raise HTTPException (status_code = 400 , detail = "content is required" )
549+ try :
550+ steps = _validate_string_list (payload .get ("steps" ), field_name = "steps" , required = True )
551+ preconditions = _validate_string_list (
552+ payload .get ("preconditions" ),
553+ field_name = "preconditions" ,
554+ required = False ,
555+ )
556+ record = ProceduralMemory (
557+ content = payload ["content" ],
558+ steps = steps ,
559+ preconditions = preconditions ,
560+ importance = float (payload .get ("importance" , 0.5 )),
561+ source = payload .get ("source" ),
562+ metadata = payload .get ("metadata" ) or {},
563+ )
564+ except ValueError as exc :
565+ raise HTTPException (status_code = 400 , detail = str (exc )) from exc
566+ service ().procedural_store .store (record )
567+ return {"record" : _serialise_record (record )}
568+
569+ @app .post ("/api/memories/procedural/file" )
570+ async def create_file_procedure (
571+ content : str = Form (...),
572+ steps : list [str ] = Form (...),
573+ preconditions : list [str ] | None = Form (default = None ),
574+ modality : str | None = Form (default = None ),
575+ media_type : str | None = Form (default = None ),
576+ text_description : str | None = Form (default = None ),
577+ importance : float = Form (default = 0.5 ),
578+ file : UploadFile = File (...),
579+ ) -> dict [str , Any ]:
580+ inferred_contract = _infer_media_contract (mime_type = file .content_type , filename = file .filename )
581+ inferred_modality = inferred_contract [0 ] if inferred_contract else None
582+ inferred_media_type = inferred_contract [1 ] if inferred_contract else None
583+ requested_media_type : str | None = None
584+ try :
585+ parsed_steps = _validate_string_list (steps , field_name = "steps" , required = True )
586+ parsed_preconditions = _validate_string_list (
587+ preconditions ,
588+ field_name = "preconditions" ,
589+ required = False ,
590+ )
591+ requested_modality = normalize_modality (modality ) if modality is not None else None
592+ resolved_modality = requested_modality or normalize_modality (inferred_modality )
593+ requested_media_type = _validate_media_type (media_type )
594+ except ValueError as exc :
595+ raise HTTPException (status_code = 400 , detail = str (exc )) from exc
596+ if requested_modality == "multimodal" and inferred_media_type is None and requested_media_type is None :
597+ raise HTTPException (
598+ status_code = 400 ,
599+ detail = "multimodal file uploads require a supported image, audio, video, or PDF file" ,
600+ )
601+ if (
602+ requested_modality is not None
603+ and inferred_modality is not None
604+ and requested_modality != inferred_modality
605+ and requested_modality != "multimodal"
606+ ):
607+ raise HTTPException (
608+ status_code = 400 ,
609+ detail = "uploaded file does not match the requested modality" ,
610+ )
611+ if resolved_modality not in _SUPPORTED_FILE_MODALITIES :
612+ raise HTTPException (
613+ status_code = 400 ,
614+ detail = "could not infer a supported modality from the uploaded file" ,
615+ )
616+ resolved_media_type = requested_media_type or inferred_media_type or resolved_modality
617+ record = ProceduralMemory (
618+ content = content ,
619+ steps = parsed_steps ,
620+ preconditions = parsed_preconditions ,
621+ modality = resolved_modality ,
622+ media_type = resolved_media_type ,
623+ text_description = text_description ,
624+ importance = importance ,
625+ )
626+ media_ref , _ = service ().save_upload (file , record .id )
627+ record .media_ref = media_ref
628+ try :
629+ service ().procedural_store .store (record )
630+ except (FileNotFoundError , ValueError ) as exc :
631+ service ().media_store .delete (media_ref )
632+ raise HTTPException (status_code = 400 , detail = str (exc )) from exc
633+ except EmbeddingProviderError as exc :
634+ service ().media_store .delete (media_ref )
635+ raise HTTPException (
636+ status_code = 502 ,
637+ detail = "Gemini embedding provider failed after retries" ,
638+ ) from exc
639+ return {"record" : _serialise_record (record )}
640+
641+ @app .post ("/api/memories/procedural/{record_id}/outcome" )
642+ async def record_procedural_outcome (record_id : str , payload : dict [str , Any ]) -> dict [str , Any ]:
643+ if "success" not in payload or not isinstance (payload ["success" ], bool ):
644+ raise HTTPException (status_code = 400 , detail = "success must be provided as a boolean" )
645+ active_service = service ()
646+ record = active_service .procedural_store .get_by_id (record_id )
647+ if record is None :
648+ raise HTTPException (status_code = 404 , detail = "procedural memory not found" )
649+ active_service .procedural_store .record_outcome (record_id , payload ["success" ])
650+ updated = active_service .procedural_store .get_by_id (record_id )
651+ return {"record" : _serialise_record (updated )}
652+
653+ @app .post ("/api/retrieval/best-procedures" )
654+ async def best_procedures (payload : dict [str , Any ]) -> dict [str , Any ]:
655+ task = payload .get ("task" , "" ).strip ()
656+ if not task :
657+ raise HTTPException (status_code = 400 , detail = "task is required" )
658+ matches = service ().procedural_store .get_best_procedure_matches (
659+ task ,
660+ top_k = int (payload .get ("top_k" , 3 )),
661+ )
662+ return {"results" : [_serialise_procedural_match (match ) for match in matches ]}
663+
492664 @app .post ("/api/retrieval/query" )
493665 async def query (payload : dict [str , Any ]) -> dict [str , Any ]:
494666 text = payload .get ("query" , "" ).strip ()
0 commit comments