Coverage for open_webui/utils/memory.py: 41%
341 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import asyncio
4import logging
5import re
6from typing import Any
8from fastapi import HTTPException
9from open_webui.models.config import Config
10from open_webui.models.memories import Memories
11from open_webui.utils.access_control import has_permission
12from open_webui.utils.json_codec import JSONCodec
13from open_webui.utils.misc import add_or_update_system_message, get_content_from_message
15log = logging.getLogger(__name__)
17MEMORY_CONTEXT_OPEN = '<memory_context>'
18MEMORY_CONTEXT_CLOSE = '</memory_context>'
21def clean_memory_content(content: str | None) -> str:
22 value = (content or '').strip()
23 if not value:
24 raise HTTPException(status_code=400, detail='Memory content cannot be empty')
25 return value
28def clean_memory_path(path: str | None) -> str | None:
29 value = re.sub(r'/+', '/', (path or '').strip().strip('/'))
30 if not value:
31 return None
32 parts = value.split('/')
33 if any(part in {'', '.', '..'} for part in parts) or any(ord(char) < 32 for char in value):
34 raise HTTPException(status_code=400, detail='Invalid memory path')
35 return value
38def memory_vector_text(content: str, path: str | None = None) -> str:
39 path = clean_memory_path(path)
40 return f'{path}\n{content}' if path else content
43def memory_label(memory) -> str:
44 return f'{memory.path}: {memory.content}' if memory.path else memory.content
47def _path_parts(path: str | None) -> list[str]:
48 return [part for part in (path or '').split('/') if part]
51def _parent_path(path: str | None) -> str | None:
52 parts = _path_parts(path)
53 return '/'.join(parts[:-1]) if len(parts) > 1 else None
56def _path_rank(memory_path: str | None, lookup_path: str | None) -> tuple | None:
57 if not lookup_path: 57 ↛ 58line 57 didn't jump to line 58 because the condition on line 57 was never true
58 return None
60 memory_path = clean_memory_path(memory_path)
61 lookup_path = clean_memory_path(lookup_path)
62 if not memory_path or not lookup_path:
63 return None
65 if memory_path == lookup_path:
66 return (0, 0)
67 if memory_path.startswith(f'{lookup_path}/'): 67 ↛ 68line 67 didn't jump to line 68 because the condition on line 67 was never true
68 return (1, len(_path_parts(memory_path)) - len(_path_parts(lookup_path)))
69 if lookup_path.startswith(f'{memory_path}/'): 69 ↛ 70line 69 didn't jump to line 70 because the condition on line 69 was never true
70 return (2, len(_path_parts(lookup_path)) - len(_path_parts(memory_path)))
71 if _parent_path(memory_path) and _parent_path(memory_path) == _parent_path(lookup_path): 71 ↛ 72line 71 didn't jump to line 72 because the condition on line 71 was never true
72 return (3, 0)
74 memory_parts = set(_path_parts(memory_path))
75 lookup_parts = set(_path_parts(lookup_path))
76 shared = len(memory_parts & lookup_parts)
77 if shared: 77 ↛ 78line 77 didn't jump to line 78 because the condition on line 77 was never true
78 return (4, -shared)
79 if _path_parts(memory_path)[-1:] == _path_parts(lookup_path)[-1:]: 79 ↛ 80line 79 didn't jump to line 80 because the condition on line 79 was never true
80 return (5, 0)
82 return None
85def _memory_matches_query(memory, query: str) -> bool:
86 value = query.strip().lower()
87 if not value:
88 return True
89 return value in (memory.content or '').lower() or value in (memory.path or '').lower()
92def search_memory_rows(
93 memories: list,
94 *,
95 query: str | None = None,
96 path: str | None = None,
97 memory_id: str | None = None,
98 memory_type: str = 'all',
99 limit: int = 20,
100) -> list:
101 rows = list(memories or [])
102 if memory_id: 102 ↛ 103line 102 didn't jump to line 103 because the condition on line 102 was never true
103 rows = [memory for memory in rows if memory.id == memory_id]
104 if memory_type != 'all':
105 rows = [memory for memory in rows if memory.type == memory_type]
107 query = (query or '').strip()
108 lookup_path = clean_memory_path(path)
109 if lookup_path:
110 basename = _path_parts(lookup_path)[-1] if _path_parts(lookup_path) else lookup_path
112 def related(memory) -> bool:
113 rank = _path_rank(memory.path, lookup_path)
114 if rank is not None:
115 return True
116 haystack = f'{memory.path or ""}\n{memory.content or ""}'.lower()
117 return lookup_path.lower() in haystack or basename.lower() in haystack
119 rows = [memory for memory in rows if related(memory)]
121 if query:
122 rows = [memory for memory in rows if _memory_matches_query(memory, query)]
124 def sort_key(memory):
125 rank = _path_rank(memory.path, lookup_path) if lookup_path else None
126 return rank if rank is not None else (9, 0), -(memory.updated_at or 0), memory.id or ''
128 return sorted(rows, key=sort_key)[: max(1, min(limit or 20, 100))]
131def list_memory_path_groups(
132 memories: list,
133 *,
134 query: str = '',
135 memory_type: str = 'all',
136 limit: int = 100,
137) -> dict:
138 rows = [
139 memory
140 for memory in (memories or [])
141 if (memory_type == 'all' or memory.type == memory_type) and _memory_matches_query(memory, query)
142 ]
143 grouped: dict[tuple[str | None, str], dict] = {}
144 for memory in rows:
145 key = (memory.path, memory.type)
146 group = grouped.setdefault(
147 key,
148 {
149 'path': memory.path,
150 'type': memory.type,
151 'count': 0,
152 'updated_at': 0,
153 'children': [],
154 },
155 )
156 group['count'] += 1
157 group['updated_at'] = max(group['updated_at'], memory.updated_at or 0)
159 paths = [path for path, _ in grouped if path]
160 for group in grouped.values():
161 path = group['path']
162 if not path:
163 continue
164 prefix = f'{path}/'
165 children = []
166 for candidate in paths:
167 if not candidate.startswith(prefix): 167 ↛ 169line 167 didn't jump to line 169 because the condition on line 167 was always true
168 continue
169 remainder = candidate[len(prefix) :]
170 child = f'{prefix}{remainder.split("/", 1)[0]}'
171 if child not in children:
172 children.append(child)
173 group['children'] = children[:20]
175 groups = sorted(grouped.values(), key=lambda item: item['updated_at'], reverse=True)
176 return {'paths': groups[: max(1, min(limit or 100, 500))], 'count': len(groups)}
179def read_memory_path_rows(
180 memories: list,
181 *,
182 path: str,
183 memory_type: str = 'all',
184 include_children: bool = True,
185 limit: int = 50,
186) -> dict:
187 lookup_path = clean_memory_path(path)
188 if not lookup_path:
189 raise HTTPException(status_code=400, detail='Memory path is required')
191 rows = [memory for memory in (memories or []) if memory_type == 'all' or memory.type == memory_type]
192 path_set = {memory.path for memory in rows if memory.path}
193 parents = [
194 '/'.join(_path_parts(lookup_path)[:idx])
195 for idx in range(1, len(_path_parts(lookup_path)))
196 if '/'.join(_path_parts(lookup_path)[:idx]) in path_set
197 ]
198 children = sorted(
199 {
200 f'{lookup_path}/{memory.path[len(lookup_path) + 1 :].split("/", 1)[0]}'
201 for memory in rows
202 if memory.path and memory.path.startswith(f'{lookup_path}/')
203 }
204 )
206 def selected(memory) -> bool:
207 if memory.path == lookup_path:
208 return True
209 if memory.path in parents: 209 ↛ 210line 209 didn't jump to line 210 because the condition on line 209 was never true
210 return True
211 return bool(include_children and memory.path and memory.path.startswith(f'{lookup_path}/'))
213 selected_rows = [memory for memory in rows if selected(memory)]
215 def sort_key(memory):
216 if memory.path == lookup_path: 216 ↛ 218line 216 didn't jump to line 218 because the condition on line 216 was always true
217 return (0, 0, -(memory.updated_at or 0), memory.id or '')
218 if memory.path and memory.path.startswith(f'{lookup_path}/'):
219 return (1, len(_path_parts(memory.path)), -(memory.updated_at or 0), memory.id or '')
220 return (2, -len(_path_parts(memory.path)), -(memory.updated_at or 0), memory.id or '')
222 return {
223 'path': lookup_path,
224 'parents': parents,
225 'children': children[:50],
226 'memories': sorted(selected_rows, key=sort_key)[: max(1, min(limit or 50, 100))],
227 }
230def memory_path_hints(query: str, memories: list, limit: int = 6) -> list[str]:
231 lowered = (query or '').lower()
232 if not lowered:
233 return []
235 hints: list[str] = []
236 for memory in sorted(memories or [], key=lambda item: (item.path or '', item.content or '', item.id or '')):
237 path = memory.path
238 if not path or path in hints:
239 continue
240 parts = _path_parts(path)
241 last = parts[-1] if parts else path
242 if path.lower() in lowered or last.lower() in lowered:
243 hints.append(path)
244 elif any(len(part) >= 3 and part.lower() in lowered for part in parts):
245 hints.append(path)
246 if len(hints) >= limit:
247 break
248 return hints
251def validate_memory_operations(form_data) -> list[dict]:
252 if not form_data.operations:
253 raise HTTPException(status_code=400, detail='No memory operations provided')
255 operations = []
256 for operation in form_data.operations:
257 op = operation.model_dump()
258 action = op.get('action')
260 if action == 'add':
261 op['content'] = clean_memory_content(op.get('content'))
262 op['type'] = Memories.normalize_memory_type(op.get('type'))
263 op['path'] = clean_memory_path(op.get('path'))
264 elif action == 'replace':
265 if not op.get('id'):
266 raise HTTPException(status_code=400, detail='Memory id is required for replace')
267 op['content'] = clean_memory_content(op.get('content'))
268 if op.get('type') is not None: 268 ↛ 270line 268 didn't jump to line 270 because the condition on line 268 was always true
269 op['type'] = Memories.normalize_memory_type(op.get('type'))
270 op['path'] = clean_memory_path(op.get('path'))
271 elif action == 'move':
272 if not op.get('id'):
273 raise HTTPException(status_code=400, detail='Memory id is required for move')
274 op['path'] = clean_memory_path(op.get('path'))
275 elif action == 'remove': 275 ↛ 279line 275 didn't jump to line 279 because the condition on line 275 was always true
276 if not op.get('id'):
277 raise HTTPException(status_code=400, detail='Memory id is required for remove')
278 else:
279 raise HTTPException(status_code=400, detail=f'Unsupported memory operation: {action}')
281 operations.append(op)
283 return operations
286def model_allows_memory(model: dict | None) -> bool:
287 return ((model or {}).get('info', {}).get('meta', {}).get('capabilities') or {}).get('memory', True)
290async def add_memory_context(request, form_data: dict, user, model: dict | None = None):
291 if not model_allows_memory(model):
292 return form_data
294 user_messages = []
295 for message in reversed(form_data.get('messages', [])):
296 if message.get('role') != 'user':
297 continue
299 content = get_content_from_message(message)
300 if isinstance(content, str) and content.strip():
301 user_messages.append(content.strip())
303 if len(user_messages) >= 7:
304 break
306 query = '\n\n'.join(reversed(user_messages))[-4000:]
307 if not query:
308 return form_data
310 all_memories = await Memories.get_memories_by_user_id(user.id)
311 results = None
312 try:
313 from open_webui.routers.memories import QueryMemoryForm, query_memory
315 results = await query_memory(request, QueryMemoryForm(content=query, k=8), user)
316 except Exception as e:
317 log.debug(e)
319 sections = {'user': [], 'neighborhood': [], 'context': []}
320 seen_ids = set()
321 for memory in sorted(
322 [memory for memory in (all_memories or []) if memory.type == 'user'],
323 key=lambda item: (item.path or '', item.updated_at or 0, item.id or ''),
324 ):
325 seen_ids.add(memory.id)
326 sections['user'].append(memory_label(memory))
328 for hint in memory_path_hints(query, all_memories):
329 for memory in search_memory_rows(
330 all_memories,
331 path=hint,
332 memory_type='context',
333 limit=4,
334 ):
335 if memory.id in seen_ids:
336 continue
337 seen_ids.add(memory.id)
338 sections['neighborhood'].append(memory_label(memory))
340 if results and hasattr(results, 'documents') and results.documents:
341 for doc_idx, doc in enumerate(results.documents[0]):
342 if not doc:
343 continue
345 metadata = {}
346 if results.metadatas and results.metadatas[0] and len(results.metadatas[0]) > doc_idx:
347 metadata = results.metadatas[0][doc_idx] or {}
349 memory_id = None
350 if results.ids and results.ids[0] and len(results.ids[0]) > doc_idx:
351 memory_id = results.ids[0][doc_idx]
352 if memory_id and memory_id in seen_ids:
353 continue
354 if memory_id:
355 seen_ids.add(memory_id)
357 content = str(doc)
358 if metadata.get('path') and content.startswith(f'{metadata.get("path")}\n'):
359 content = content[len(metadata.get('path')) + 1 :]
360 label = f'{metadata.get("path")}: {content}' if metadata.get('path') else content
361 sections[Memories.normalize_memory_type(metadata.get('type'))].append(label)
363 parts = []
364 for title, key in (
365 ('User Memory', 'user'),
366 ('Memory Neighborhood', 'neighborhood'),
367 ('Relevant Context', 'context'),
368 ):
369 if sections[key]:
370 ordered = sorted(sections[key], key=lambda memory: (memory.casefold(), memory))
371 parts.append(f'[{title}]\n' + '\n'.join(f'- {memory}' for memory in ordered))
372 if not parts:
373 return form_data
375 config = await Config.get_many('memories.user_char_limit', 'memories.context_char_limit')
376 try:
377 user_limit = max(250, int(config.get('memories.user_char_limit') or 2000))
378 except Exception:
379 user_limit = 2000
380 try:
381 context_limit = max(250, int(config.get('memories.context_char_limit') or 2000))
382 except Exception:
383 context_limit = 2000
385 messages = form_data['messages']
386 if messages and messages[0].get('role') == 'system':
387 content = messages[0].get('content', '')
388 if isinstance(content, str) and MEMORY_CONTEXT_OPEN in content:
389 start = content.find(MEMORY_CONTEXT_OPEN)
390 end = content.find(MEMORY_CONTEXT_CLOSE, start)
391 if end != -1:
392 messages[0]['content'] = (content[:start] + content[end + len(MEMORY_CONTEXT_CLOSE) :]).strip()
394 user_parts = [part for part in parts if part.startswith('[User Memory]')]
395 context_parts = [part for part in parts if not part.startswith('[User Memory]')]
396 rendered = '\n\n'.join(
397 [
398 '\n\n'.join(user_parts)[:user_limit],
399 '\n\n'.join(context_parts)[:context_limit],
400 ]
401 ).strip()
402 if not rendered:
403 return form_data
405 memory_context = f'{MEMORY_CONTEXT_OPEN}\n{rendered}\n{MEMORY_CONTEXT_CLOSE}'
406 form_data['messages'] = add_or_update_system_message(memory_context, messages, append=True)
407 return form_data
410async def review_memory_after_turn(
411 *,
412 request,
413 user,
414 model: dict | None,
415 metadata: dict,
416 form_data: dict,
417 assistant_message: dict,
418 messages: list[dict],
419) -> None:
420 if not model_allows_memory(model):
421 return
423 features = metadata.get('features') or {}
424 if not features.get('memory'):
425 return
427 assistant_content = get_content_from_message(assistant_message)
428 if not isinstance(assistant_content, str) or not assistant_content.strip():
429 return
431 config = await Config.get_many(
432 'memories.enable',
433 'memories.background_review.enable',
434 'memories.review_interval_turns',
435 'user.permissions',
436 )
437 if not config.get('memories.enable') or not config.get('memories.background_review.enable'):
438 return
440 try:
441 interval = max(1, int(config.get('memories.review_interval_turns', 10)))
442 except Exception:
443 interval = 10
445 user_turns = len([message for message in messages if message.get('role') == 'user'])
446 if user_turns == 0 or user_turns % interval != 0:
447 return
449 # features is client-supplied; re-check the permission the memory routes enforce.
450 if user.role != 'admin' and not await has_permission(user.id, 'features.memories', config.get('user.permissions')):
451 return
453 task = asyncio.create_task(
454 _review_memory(
455 request=request,
456 user=user,
457 model=model,
458 metadata=metadata,
459 form_data=form_data,
460 assistant_message=assistant_message,
461 messages=messages,
462 )
463 )
465 def log_failure(done_task):
466 try:
467 done_task.result()
468 except Exception as e:
469 log.debug('Memory review failed: %s', e)
471 task.add_done_callback(log_failure)
474async def _review_memory(
475 *,
476 request,
477 user,
478 model: dict | None,
479 metadata: dict,
480 form_data: dict,
481 assistant_message: dict,
482 messages: list[dict],
483) -> None:
484 existing_memories = await Memories.get_memories_by_user_id(user.id)
485 existing_lines = [
486 f'- id={memory.id} type={memory.type} path={memory.path or ""} content={memory.content}'
487 for memory in (existing_memories or [])[:80]
488 ]
490 assistant_content = get_content_from_message(assistant_message)
491 if not isinstance(assistant_content, str):
492 assistant_content = ''
494 transcript_lines = []
495 for message in messages[-16:]:
496 role = message.get('role', '')
497 content = message.get('content', '')
498 if not isinstance(content, str):
499 content = get_content_from_message(message)
500 content = content.strip()
501 if role not in {'user', 'assistant'} or not content:
502 continue
503 if len(content) > 1600:
504 content = f'{content[:1000]}\n...(truncated)...\n{content[-400:]}'
505 transcript_lines.append(f'{role}: {content}')
507 if assistant_content.strip():
508 assistant_final = assistant_content.strip()
509 if len(assistant_final) > 1600:
510 assistant_final = f'{assistant_final[:1000]}\n...(truncated)...\n{assistant_final[-400:]}'
511 transcript_lines.append(f'assistant_final: {assistant_final}')
513 model_id = model.get('id') if isinstance(model, dict) else form_data.get('model')
514 operations = await _generate_memory_operations(
515 request=request,
516 user=user,
517 model_id=model_id,
518 metadata=metadata,
519 existing_text='\n'.join(existing_lines) if existing_lines else '(none)',
520 transcript='\n\n'.join(transcript_lines),
521 )
522 if operations:
523 from open_webui.routers.memories import UpdateMemoriesForm, update_memories
525 await update_memories(request, UpdateMemoriesForm(operations=operations, source='background_review'), user)
528async def _generate_memory_operations(
529 *,
530 request,
531 user,
532 model_id: str,
533 metadata: dict,
534 existing_text: str,
535 transcript: str,
536) -> list[dict[str, Any]]:
537 from open_webui.utils.chat import generate_chat_completion
539 review_prompt = f"""Review the completed conversation turn and decide whether long-term memory should change.
541Memory types:
542- user: durable facts, preferences, or instructions about the user.
543- context: other durable context that may help future chats for this user account.
545Rules:
546- Save enduring details that can improve future conversations.
547- Do not save one-off activity, meals, temporary mood, routine daily events, or other short-lived details unless the user explicitly asks to remember them.
548- Do not save secrets, credentials, transient task steps, or unsupported guesses.
549- Use path when there is a clear path for the memory.
550- Leave path empty when there is no clear place for the memory.
551- Prefer replace/move/remove over duplicate add when an existing memory should change.
552- Do not invent type, status, trait, score, importance, or stability schemas.
553- Return only JSON in this shape:
554 {{"operations":[
555 {{"action":"add","type":"user|context","path":"...","content":"..."}},
556 {{"action":"replace","id":"...","type":"user|context","path":"...","content":"..."}},
557 {{"action":"move","id":"...","path":"..."}},
558 {{"action":"remove","id":"..."}}
559 ]}}
560- Use an empty operations array if nothing should be remembered.
562Existing memories:
563{existing_text}
565Conversation:
566{transcript}
567"""
569 response = await generate_chat_completion(
570 request,
571 form_data={
572 'model': model_id,
573 'messages': [
574 {
575 'role': 'system',
576 # LICENSE covers this Open WebUI system identifier.
577 # Do not alter, remove, obscure, or replace it except as LICENSE permits:
578 # https://docs.openwebui.com/license.
579 'content': "You are Open WebUI's private memory reviewer. Return only valid JSON.",
580 },
581 {'role': 'user', 'content': review_prompt},
582 ],
583 'stream': False,
584 'metadata': {
585 'task': 'memory_review',
586 'chat_id': metadata.get('chat_id'),
587 'message_id': metadata.get('message_id'),
588 },
589 },
590 user=user,
591 )
593 if not isinstance(response, dict) or not response.get('choices'):
594 return []
596 response_message = response.get('choices', [{}])[0].get('message', {})
597 content = response_message.get('content') or response_message.get('reasoning_content') or ''
598 start = content.find('{')
599 end = content.rfind('}')
600 if start == -1 or end == -1 or end < start:
601 return []
603 try:
604 parsed = JSONCodec.loads(content[start : end + 1])
605 except Exception:
606 return []
608 operations = parsed.get('operations') if isinstance(parsed, dict) else None
609 return operations if isinstance(operations, list) else []