Skip to content

Commit 28eb916

Browse files
authored
feat(api): support bulk memory associations (#198)
## Summary - Add backwards-compatible batch support to POST /associate via an associations array capped at 500 items - Batch-write valid associations with FalkorDB UNWIND grouped by relationship type and return per-index partial-success results - Update the MCP bridge, tests, and docs for bulk associate_memories behavior Closes #196. ## Tests - make test - make lint - source .venv/bin/activate && black --check automem/api/memory.py tests/support/fake_graph.py tests/test_api_endpoints.py - npm test --prefix mcp-sse-server ## API / Config - Existing single-association request/response shape remains supported - Batch responses return 201 when all items succeed and 207 when any item fails - No config changes
2 parents 431433e + ea4e08f commit 28eb916

7 files changed

Lines changed: 789 additions & 18 deletions

File tree

‎automem/api/memory.py‎

Lines changed: 249 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,236 @@ def _validate_memory_id(memory_id: str) -> None:
2626
abort(400, description="memory_id must be a valid UUID")
2727

2828

29+
def _uuid_error(memory_id: str, field_name: str) -> Optional[str]:
30+
try:
31+
uuid.UUID(memory_id)
32+
except ValueError:
33+
return f"'{field_name}' must be a valid UUID"
34+
return None
35+
36+
37+
def _association_summary(created_count: int, total_count: int) -> str:
38+
return f"{created_count}/{total_count} associations created successfully"
39+
40+
41+
def _association_failure(index: int, reason: str) -> Dict[str, Any]:
42+
return {"index": index, "reason": reason}
43+
44+
45+
def _association_success(
46+
*,
47+
index: int,
48+
memory1_id: str,
49+
memory2_id: str,
50+
relation_type: str,
51+
strength: float,
52+
) -> Dict[str, Any]:
53+
return {
54+
"index": index,
55+
"memory1_id": memory1_id,
56+
"memory2_id": memory2_id,
57+
"relation_type": relation_type,
58+
"strength": strength,
59+
}
60+
61+
62+
def _prepare_association_props(
63+
*,
64+
payload: Dict[str, Any],
65+
relation_type: str,
66+
strength: float,
67+
timestamp: str,
68+
relationship_types: Dict[str, Dict[str, Any]],
69+
) -> Dict[str, Any]:
70+
relationship_props = {"strength": strength, "updated_at": timestamp}
71+
relation_config = relationship_types.get(relation_type, {})
72+
for prop in relation_config.get("properties", []):
73+
if prop in payload and prop not in relationship_props:
74+
relationship_props[prop] = payload[prop]
75+
return relationship_props
76+
77+
78+
def _parse_association_item(
79+
*,
80+
item: Any,
81+
index: int,
82+
authorable_relations: Set[str],
83+
relationship_types: Dict[str, Dict[str, Any]],
84+
coerce_importance_fn: Callable[[Any], float],
85+
timestamp: str,
86+
) -> tuple[Optional[Dict[str, Any]], Optional[Dict[str, Any]]]:
87+
if not isinstance(item, dict):
88+
return None, _association_failure(index, "Association item must be an object")
89+
90+
memory1_id = str(item.get("memory1_id") or "").strip()
91+
memory2_id = str(item.get("memory2_id") or "").strip()
92+
relation_type = str(item.get("type") or "RELATES_TO").strip().upper()
93+
strength = coerce_importance_fn(item.get("strength", 0.5))
94+
95+
if not memory1_id or not memory2_id:
96+
return None, _association_failure(index, "'memory1_id' and 'memory2_id' are required")
97+
98+
for field_name, value in (("memory1_id", memory1_id), ("memory2_id", memory2_id)):
99+
error = _uuid_error(value, field_name)
100+
if error:
101+
return None, _association_failure(index, error)
102+
103+
if memory1_id == memory2_id:
104+
return None, _association_failure(index, "Cannot associate a memory with itself")
105+
106+
if relation_type not in authorable_relations:
107+
return None, _association_failure(
108+
index,
109+
f"Relation type must be one of {sorted(authorable_relations)}",
110+
)
111+
112+
return (
113+
{
114+
"index": index,
115+
"memory1_id": memory1_id,
116+
"memory2_id": memory2_id,
117+
"type": relation_type,
118+
"strength": strength,
119+
"props": _prepare_association_props(
120+
payload=item,
121+
relation_type=relation_type,
122+
strength=strength,
123+
timestamp=timestamp,
124+
relationship_types=relationship_types,
125+
),
126+
},
127+
None,
128+
)
129+
130+
131+
def _batch_association_response(
132+
*,
133+
succeeded: List[Dict[str, Any]],
134+
failed: List[Dict[str, Any]],
135+
total_count: int,
136+
jsonify_fn: Callable[[Any], Any],
137+
) -> Any:
138+
created_count = len(succeeded)
139+
failed_count = len(failed)
140+
status_code = 201 if failed_count == 0 else 207
141+
status = "success" if failed_count == 0 else "partial_success"
142+
return (
143+
jsonify_fn(
144+
{
145+
"status": status,
146+
"created_count": created_count,
147+
"failed_count": failed_count,
148+
"succeeded": sorted(succeeded, key=lambda item: item["index"]),
149+
"failed": sorted(failed, key=lambda item: item["index"]),
150+
"summary": _association_summary(created_count, total_count),
151+
}
152+
),
153+
status_code,
154+
)
155+
156+
157+
def _create_association_batch(
158+
*,
159+
payload: Dict[str, Any],
160+
coerce_importance_fn: Callable[[Any], float],
161+
get_memory_graph_fn: Callable[[], Any],
162+
authorable_relations: Set[str],
163+
relationship_types: Dict[str, Dict[str, Any]],
164+
utc_now_fn: Callable[[], str],
165+
abort_fn: Callable[..., Any],
166+
jsonify_fn: Callable[[Any], Any],
167+
logger: Any,
168+
) -> Any:
169+
associations = payload.get("associations")
170+
if not isinstance(associations, list) or len(associations) == 0:
171+
abort_fn(400, description="'associations' must be a non-empty array")
172+
if len(associations) > 500:
173+
abort_fn(400, description="Batch size limit is 500 associations per request")
174+
175+
timestamp = utc_now_fn()
176+
valid_rows: List[Dict[str, Any]] = []
177+
failed: List[Dict[str, Any]] = []
178+
for index, item in enumerate(associations):
179+
row, failure = _parse_association_item(
180+
item=item,
181+
index=index,
182+
authorable_relations=authorable_relations,
183+
relationship_types=relationship_types,
184+
coerce_importance_fn=coerce_importance_fn,
185+
timestamp=timestamp,
186+
)
187+
if failure:
188+
failed.append(failure)
189+
elif row:
190+
valid_rows.append(row)
191+
192+
succeeded: List[Dict[str, Any]] = []
193+
if valid_rows:
194+
graph = get_memory_graph_fn()
195+
if graph is None:
196+
abort_fn(503, description="FalkorDB is unavailable")
197+
198+
rows_by_index = {row["index"]: row for row in valid_rows}
199+
rows_by_type: Dict[str, List[Dict[str, Any]]] = {}
200+
for row in valid_rows:
201+
rows_by_type.setdefault(row["type"], []).append(row)
202+
203+
for relation_type, rows in rows_by_type.items():
204+
try:
205+
result = graph.query(
206+
f"""
207+
UNWIND $rows AS row
208+
MATCH (m1:Memory {{id: row.memory1_id}})
209+
MATCH (m2:Memory {{id: row.memory2_id}})
210+
MERGE (m1)-[r:{relation_type}]->(m2)
211+
SET r += row.props
212+
RETURN row.index, row.memory1_id, row.memory2_id
213+
""",
214+
{"rows": rows},
215+
)
216+
except Exception:
217+
logger.exception(
218+
"Failed to create association batch for relation type %s",
219+
relation_type,
220+
)
221+
for row in rows:
222+
failed.append(
223+
_association_failure(
224+
row["index"],
225+
f"Failed to create association batch for relation type {relation_type}",
226+
)
227+
)
228+
continue
229+
230+
created_indexes = set()
231+
for result_row in list(getattr(result, "result_set", []) or []):
232+
index = int(result_row[0])
233+
created_indexes.add(index)
234+
source = rows_by_index[index]
235+
succeeded.append(
236+
_association_success(
237+
index=index,
238+
memory1_id=source["memory1_id"],
239+
memory2_id=source["memory2_id"],
240+
relation_type=source["type"],
241+
strength=source["strength"],
242+
)
243+
)
244+
245+
for row in rows:
246+
if row["index"] not in created_indexes:
247+
failed.append(
248+
_association_failure(row["index"], "One or both memories do not exist")
249+
)
250+
251+
return _batch_association_response(
252+
succeeded=succeeded,
253+
failed=failed,
254+
total_count=len(associations),
255+
jsonify_fn=jsonify_fn,
256+
)
257+
258+
29259
def _parse_by_tag_request(
30260
*,
31261
request_args: Any,
@@ -767,6 +997,18 @@ def associate() -> Any:
767997
payload = request.get_json(silent=True)
768998
if not isinstance(payload, dict):
769999
abort(400, description="JSON body is required")
1000+
if isinstance(payload.get("associations"), list):
1001+
return _create_association_batch(
1002+
payload=payload,
1003+
coerce_importance_fn=coerce_importance,
1004+
get_memory_graph_fn=get_memory_graph,
1005+
authorable_relations=set(authorable_relations),
1006+
relationship_types=relation_types,
1007+
utc_now_fn=utc_now,
1008+
abort_fn=abort,
1009+
jsonify_fn=jsonify,
1010+
logger=logger,
1011+
)
7701012

7711013
memory1_id = (payload.get("memory1_id") or "").strip()
7721014
memory2_id = (payload.get("memory2_id") or "").strip()
@@ -791,12 +1033,14 @@ def associate() -> Any:
7911033

7921034
timestamp = utc_now()
7931035

794-
relationship_props = {"strength": strength, "updated_at": timestamp}
1036+
relationship_props = _prepare_association_props(
1037+
payload=payload,
1038+
relation_type=relation_type,
1039+
strength=strength,
1040+
timestamp=timestamp,
1041+
relationship_types=relation_types,
1042+
)
7951043
relation_config = relation_types.get(relation_type, {})
796-
if "properties" in relation_config:
797-
for prop in relation_config["properties"]:
798-
if prop in payload:
799-
relationship_props[prop] = payload[prop]
8001044

8011045
set_clauses = [f"r.{key} = ${key}" for key in relationship_props]
8021046
set_clause = ", ".join(set_clauses)

‎docs/API.md‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,7 +41,10 @@ Memory
4141

4242
- POST `/associate`
4343
- Body: `{ "memory1_id": "...", "memory2_id": "...", "type": "RELATES_TO", "strength": 0.9 }`
44-
- Response: `{ "status": "success", ... }`
44+
- Batch body: `{ "associations": [{ "memory1_id": "...", "memory2_id": "...", "type": "RELATES_TO", "strength": 0.9 }] }` (max 500)
45+
- Single response: `{ "status": "success", ... }`
46+
- Batch response: `{ "status": "success"|"partial_success", "created_count": C, "failed_count": F, "succeeded": [...], "failed": [...], "summary": "C/M associations created successfully" }`
47+
- Batch requests return `201` when every item succeeds and `207` when one or more items fail validation or reference missing memories.
4548
- `type` accepts only the 11 authorable semantic relationship types. System-generated labels such as `SIMILAR_TO`, `PRECEDED_BY`, and `DISCOVERED` are readable/filterable but cannot be created via this endpoint.
4649

4750
Recall

‎docs/MCP_SSE.md‎

Lines changed: 27 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -155,7 +155,33 @@ The MCP bridge exposes these MCP tools:
155155

156156
## Recall Ordering
157157

158-
`associate_memories` accepts only the 11 authorable semantic relationship types. Auto-generated labels such as `SIMILAR_TO`, `PRECEDED_BY`, and `DISCOVERED` remain readable on recall surfaces but are not exposed as authoring choices in the MCP schema.
158+
`associate_memories` accepts either a single association:
159+
160+
```json
161+
{
162+
"memory1_id": "...",
163+
"memory2_id": "...",
164+
"type": "RELATES_TO",
165+
"strength": 0.9
166+
}
167+
```
168+
169+
or a batch of up to 500 associations:
170+
171+
```json
172+
{
173+
"associations": [
174+
{
175+
"memory1_id": "...",
176+
"memory2_id": "...",
177+
"type": "RELATES_TO",
178+
"strength": 0.9
179+
}
180+
]
181+
}
182+
```
183+
184+
Batch calls return a concise summary and item-indexed failures when only some associations succeed. `associate_memories` accepts only the 11 authorable semantic relationship types. Auto-generated labels such as `SIMILAR_TO`, `PRECEDED_BY`, and `DISCOVERED` remain readable on recall surfaces but are not exposed as authoring choices in the MCP schema.
159185

160186
`recall_memory` defaults to relevance ranking (`sort: "score"`). For chronological recaps (e.g. “what happened since X”), set:
161187

0 commit comments

Comments
 (0)