Skip to content

Commit b3b41ca

Browse files
Add procedural memory store and surfaces
1 parent 89a669f commit b3b41ca

8 files changed

Lines changed: 1249 additions & 15 deletions

File tree

api/app.py

Lines changed: 174 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,10 +18,12 @@
1818
from events import EventBus
1919
from models.base import MemoryRecord, normalize_modality
2020
from models.episodic import EpisodicMemory
21+
from models.procedural import ProceduralMemory
2122
from models.semantic import SemanticMemory
2223
from retrieval.retriever import UnifiedRetriever
2324
from stores.episodic_store import EpisodicStore, EpisodicStoreError, MediaTooLargeError
2425
from stores.media_store import MediaStore
26+
from stores.procedural_store import ProceduralMatch, ProceduralStore
2527
from stores.semantic_store import SemanticStore
2628
from 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+
144165
def _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+
231273
class 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

Comments
 (0)