mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-28 05:27:35 +00:00
refac
This commit is contained in:
parent
7e13fd7ad1
commit
dbdcfd8c60
5 changed files with 507 additions and 70 deletions
|
|
@ -2,7 +2,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Optional
|
||||
from typing import Literal, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
|
|
@ -14,6 +14,7 @@ from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT
|
|||
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
|
||||
from open_webui.utils.access_control import has_permission
|
||||
from open_webui.utils.auth import get_verified_user
|
||||
from open_webui.utils.memory import clean_memory_content, validate_memory_operations
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
|
|
@ -64,10 +65,23 @@ async def get_memories(
|
|||
|
||||
class AddMemoryForm(BaseModel):
|
||||
content: str
|
||||
type: Literal['user', 'context'] = 'context'
|
||||
|
||||
|
||||
class MemoryUpdateModel(BaseModel):
|
||||
content: str | None = None
|
||||
type: Literal['user', 'context'] | None = None
|
||||
|
||||
|
||||
class MemoryOperationModel(BaseModel):
|
||||
action: Literal['add', 'replace', 'remove']
|
||||
id: str | None = None
|
||||
content: str | None = None
|
||||
type: Literal['user', 'context'] | None = None
|
||||
|
||||
|
||||
class UpdateMemoriesForm(BaseModel):
|
||||
operations: list[MemoryOperationModel]
|
||||
|
||||
|
||||
@router.post('/add', response_model=MemoryModel | None)
|
||||
|
|
@ -84,7 +98,8 @@ async def add_memory(
|
|||
"""
|
||||
await check_memories_permission(user)
|
||||
|
||||
memory = await Memories.insert_new_memory(user.id, form_data.content)
|
||||
content = clean_memory_content(form_data.content)
|
||||
memory = await Memories.insert_new_memory(user.id, content, memory_type=form_data.type)
|
||||
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
|
||||
|
|
@ -95,7 +110,11 @@ async def add_memory(
|
|||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'vector': vector,
|
||||
'metadata': {'created_at': memory.created_at},
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
'type': memory.type,
|
||||
},
|
||||
}
|
||||
],
|
||||
)
|
||||
|
|
@ -105,11 +124,87 @@ async def add_memory(
|
|||
EVENTS.MEMORY_CREATED,
|
||||
actor=user,
|
||||
subject_id=memory.id,
|
||||
data={'content_preview': memory.content[:300]},
|
||||
data={'content_preview': memory.content[:300], 'type': memory.type},
|
||||
)
|
||||
return memory
|
||||
|
||||
|
||||
@router.post('/update', response_model=list[dict])
|
||||
async def update_memories(
|
||||
request: Request,
|
||||
form_data: UpdateMemoriesForm,
|
||||
user=Depends(get_verified_user),
|
||||
):
|
||||
await check_memories_permission(user)
|
||||
|
||||
operations = validate_memory_operations(form_data)
|
||||
|
||||
try:
|
||||
results = await Memories.apply_memory_operations(user.id, operations)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
|
||||
upsert_items = []
|
||||
delete_ids = []
|
||||
response = []
|
||||
|
||||
for result in results:
|
||||
memory = result.get('memory')
|
||||
if isinstance(memory, MemoryModel):
|
||||
result = {**result, 'memory': memory.model_dump()}
|
||||
if result.get('status') in {'created', 'updated'}:
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
upsert_items.append(
|
||||
{
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'vector': vector,
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
'type': memory.type,
|
||||
},
|
||||
}
|
||||
)
|
||||
if result.get('status') == 'deleted' and result.get('id'):
|
||||
delete_ids.append(result['id'])
|
||||
response.append(result)
|
||||
|
||||
if upsert_items:
|
||||
await ASYNC_VECTOR_DB_CLIENT.upsert(collection_name=f'user-memory-{user.id}', items=upsert_items)
|
||||
|
||||
if delete_ids:
|
||||
await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=delete_ids)
|
||||
|
||||
for result in response:
|
||||
status_value = result.get('status')
|
||||
memory = result.get('memory') or {}
|
||||
memory_id = memory.get('id') or result.get('id')
|
||||
|
||||
if status_value == 'created':
|
||||
event = EVENTS.MEMORY_CREATED
|
||||
elif status_value == 'updated':
|
||||
event = EVENTS.MEMORY_UPDATED
|
||||
elif status_value == 'deleted':
|
||||
event = EVENTS.MEMORY_DELETED
|
||||
else:
|
||||
continue
|
||||
|
||||
await publish_event(
|
||||
request,
|
||||
event,
|
||||
actor=user,
|
||||
subject_id=memory_id,
|
||||
data={
|
||||
'content_preview': (memory.get('content') or '')[:300],
|
||||
'type': memory.get('type'),
|
||||
'operation': result.get('action'),
|
||||
},
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
############################
|
||||
# QueryMemory
|
||||
############################
|
||||
|
|
@ -216,6 +311,7 @@ async def reset_memory_from_vector_db(
|
|||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
'type': memory.type,
|
||||
},
|
||||
}
|
||||
for idx, memory in enumerate(memories)
|
||||
|
|
@ -281,7 +377,10 @@ async def update_memory_by_id(
|
|||
# EMBEDDING_FUNCTION() which makes external API calls (1-5+ seconds).
|
||||
await check_memories_permission(user)
|
||||
|
||||
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content)
|
||||
content = clean_memory_content(form_data.content) if form_data.content is not None else None
|
||||
if content is None and form_data.type is None:
|
||||
raise HTTPException(status_code=400, detail='No memory update provided')
|
||||
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, content, memory_type=form_data.type)
|
||||
if memory is None:
|
||||
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
|
|
@ -298,6 +397,7 @@ async def update_memory_by_id(
|
|||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
'type': memory.type,
|
||||
},
|
||||
}
|
||||
],
|
||||
|
|
@ -308,7 +408,7 @@ async def update_memory_by_id(
|
|||
EVENTS.MEMORY_UPDATED,
|
||||
actor=user,
|
||||
subject_id=memory.id,
|
||||
data={'content_preview': memory.content[:300]},
|
||||
data={'content_preview': memory.content[:300], 'type': memory.type},
|
||||
)
|
||||
return memory
|
||||
|
||||
|
|
|
|||
|
|
@ -36,7 +36,9 @@ from open_webui.routers.memories import (
|
|||
AddMemoryForm,
|
||||
MemoryUpdateModel,
|
||||
QueryMemoryForm,
|
||||
UpdateMemoriesForm,
|
||||
query_memory,
|
||||
update_memories as _update_memories,
|
||||
update_memory_by_id,
|
||||
)
|
||||
from open_webui.routers.memories import (
|
||||
|
|
@ -609,28 +611,41 @@ async def search_memories(
|
|||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
|
||||
results = await query_memory(
|
||||
memory_results = await query_memory(
|
||||
__request__,
|
||||
QueryMemoryForm(content=query, k=count),
|
||||
user,
|
||||
)
|
||||
|
||||
if results and hasattr(results, 'documents') and results.documents:
|
||||
memories = []
|
||||
for doc_idx, doc in enumerate(results.documents[0]):
|
||||
memory_id = None
|
||||
if results.ids and results.ids[0]:
|
||||
memory_id = results.ids[0][doc_idx]
|
||||
created_at = 'Unknown'
|
||||
if results.metadatas and results.metadatas[0][doc_idx].get('created_at'):
|
||||
created_at = time.strftime(
|
||||
'%Y-%m-%d',
|
||||
time.localtime(results.metadatas[0][doc_idx]['created_at']),
|
||||
)
|
||||
memories.append({'id': memory_id, 'date': created_at, 'content': doc})
|
||||
return json.dumps(memories, ensure_ascii=False)
|
||||
else:
|
||||
if not memory_results or not hasattr(memory_results, 'documents') or not memory_results.documents:
|
||||
return json.dumps([])
|
||||
|
||||
memories = []
|
||||
for memory_index, memory_text in enumerate(memory_results.documents[0]):
|
||||
memory_id = None
|
||||
if memory_results.ids and memory_results.ids[0]:
|
||||
memory_id = memory_results.ids[0][memory_index]
|
||||
|
||||
metadata = {}
|
||||
if memory_results.metadatas and memory_results.metadatas[0] and len(memory_results.metadatas[0]) > memory_index:
|
||||
metadata = memory_results.metadatas[0][memory_index] or {}
|
||||
|
||||
created_at = 'Unknown'
|
||||
if metadata.get('created_at'):
|
||||
created_at = time.strftime(
|
||||
'%Y-%m-%d',
|
||||
time.localtime(metadata['created_at']),
|
||||
)
|
||||
|
||||
memories.append(
|
||||
{
|
||||
'id': memory_id,
|
||||
'type': Memories.normalize_memory_type(metadata.get('type')),
|
||||
'date': created_at,
|
||||
'content': memory_text,
|
||||
}
|
||||
)
|
||||
return json.dumps(memories, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'search_memories error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
|
@ -638,6 +653,7 @@ async def search_memories(
|
|||
|
||||
async def add_memory(
|
||||
content: str,
|
||||
type: str = 'user',
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
|
|
@ -645,6 +661,7 @@ async def add_memory(
|
|||
Save a user-provided preference, fact, or instruction as memory for future chats.
|
||||
|
||||
:param content: The memory content to store
|
||||
:param type: Use "user" for facts/preferences about the user, or "context" for other durable context
|
||||
:return: Confirmation that the memory was stored
|
||||
"""
|
||||
if __request__ is None:
|
||||
|
|
@ -655,19 +672,55 @@ async def add_memory(
|
|||
|
||||
memory = await _add_memory(
|
||||
__request__,
|
||||
AddMemoryForm(content=content),
|
||||
AddMemoryForm(content=content, type=Memories.normalize_memory_type(type)),
|
||||
user,
|
||||
)
|
||||
|
||||
return json.dumps({'status': 'success', 'id': memory.id}, ensure_ascii=False)
|
||||
return json.dumps({'status': 'success', 'id': memory.id, 'type': memory.type}, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'add_memory error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def update_memory(
|
||||
operations: list[dict],
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
"""
|
||||
Apply a batch of memory changes after learning durable information.
|
||||
|
||||
Use type "user" for facts, preferences, or instructions about the user.
|
||||
Use type "context" for other durable context that may help future chats.
|
||||
|
||||
Operation shapes:
|
||||
- {"action": "add", "content": "...", "type": "user"|"context"}
|
||||
- {"action": "replace", "id": "...", "content": "...", "type": "user"|"context"}
|
||||
- {"action": "remove", "id": "..."}
|
||||
|
||||
:param operations: Memory operations to apply in one request
|
||||
:return: JSON with operation results
|
||||
"""
|
||||
if __request__ is None:
|
||||
return json.dumps({'error': 'Request context not available'})
|
||||
|
||||
try:
|
||||
user = UserModel(**__user__) if __user__ else None
|
||||
operation_results = await _update_memories(
|
||||
__request__,
|
||||
UpdateMemoriesForm(operations=operations),
|
||||
user,
|
||||
)
|
||||
return json.dumps(operation_results, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
log.exception(f'update_memory error: {e}')
|
||||
return json.dumps({'error': str(e)})
|
||||
|
||||
|
||||
async def replace_memory_content(
|
||||
memory_id: str,
|
||||
content: str,
|
||||
type: Optional[str] = None,
|
||||
__request__: Request = None,
|
||||
__user__: dict = None,
|
||||
) -> str:
|
||||
|
|
@ -676,6 +729,7 @@ async def replace_memory_content(
|
|||
|
||||
:param memory_id: The ID of the memory to update
|
||||
:param content: The new content for the memory
|
||||
:param type: Optional "user" or "context" type for the updated memory
|
||||
:return: Confirmation that the memory was updated
|
||||
"""
|
||||
if __request__ is None:
|
||||
|
|
@ -687,12 +741,12 @@ async def replace_memory_content(
|
|||
memory = await update_memory_by_id(
|
||||
memory_id=memory_id,
|
||||
request=__request__,
|
||||
form_data=MemoryUpdateModel(content=content),
|
||||
form_data=MemoryUpdateModel(content=content, type=Memories.normalize_memory_type(type) if type else None),
|
||||
user=user,
|
||||
)
|
||||
|
||||
return json.dumps(
|
||||
{'status': 'success', 'id': memory.id, 'content': memory.content},
|
||||
{'status': 'success', 'id': memory.id, 'type': memory.type, 'content': memory.content},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -750,16 +804,17 @@ async def list_memories(
|
|||
memories = await Memories.get_memories_by_user_id(user.id)
|
||||
|
||||
if memories:
|
||||
result = [
|
||||
memory_rows = [
|
||||
{
|
||||
'id': m.id,
|
||||
'type': m.type,
|
||||
'content': m.content,
|
||||
'created_at': time.strftime('%Y-%m-%d %H:%M', time.localtime(m.created_at)),
|
||||
'updated_at': time.strftime('%Y-%m-%d %H:%M', time.localtime(m.updated_at)),
|
||||
}
|
||||
for m in memories
|
||||
]
|
||||
return json.dumps(result, ensure_ascii=False)
|
||||
return json.dumps(memory_rows, ensure_ascii=False)
|
||||
else:
|
||||
return json.dumps([])
|
||||
except Exception as e:
|
||||
|
|
|
|||
314
backend/open_webui/utils/memory.py
Normal file
314
backend/open_webui/utils/memory.py
Normal file
|
|
@ -0,0 +1,314 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from open_webui.models.config import Config
|
||||
from open_webui.models.memories import Memories
|
||||
from open_webui.utils.misc import add_or_update_system_message, get_content_from_message
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
MEMORY_CONTEXT_OPEN = '<memory_context>'
|
||||
MEMORY_CONTEXT_CLOSE = '</memory_context>'
|
||||
|
||||
|
||||
def clean_memory_content(content: str | None) -> str:
|
||||
value = (content or '').strip()
|
||||
if not value:
|
||||
raise HTTPException(status_code=400, detail='Memory content cannot be empty')
|
||||
return value
|
||||
|
||||
|
||||
def validate_memory_operations(form_data) -> list[dict]:
|
||||
if not form_data.operations:
|
||||
raise HTTPException(status_code=400, detail='No memory operations provided')
|
||||
|
||||
operations = []
|
||||
for operation in form_data.operations:
|
||||
op = operation.model_dump()
|
||||
action = op.get('action')
|
||||
|
||||
if action == 'add':
|
||||
op['content'] = clean_memory_content(op.get('content'))
|
||||
op['type'] = Memories.normalize_memory_type(op.get('type'))
|
||||
elif action == 'replace':
|
||||
if not op.get('id'):
|
||||
raise HTTPException(status_code=400, detail='Memory id is required for replace')
|
||||
op['content'] = clean_memory_content(op.get('content'))
|
||||
if op.get('type') is not None:
|
||||
op['type'] = Memories.normalize_memory_type(op.get('type'))
|
||||
elif action == 'remove':
|
||||
if not op.get('id'):
|
||||
raise HTTPException(status_code=400, detail='Memory id is required for remove')
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=f'Unsupported memory operation: {action}')
|
||||
|
||||
operations.append(op)
|
||||
|
||||
return operations
|
||||
|
||||
|
||||
def model_allows_memory(model: dict | None) -> bool:
|
||||
return (model or {}).get('info', {}).get('meta', {}).get('capabilities', {}).get('memory', True)
|
||||
|
||||
|
||||
async def add_memory_context(request, form_data: dict, user, model: dict | None = None):
|
||||
if not model_allows_memory(model):
|
||||
return form_data
|
||||
|
||||
user_messages = []
|
||||
for message in reversed(form_data.get('messages', [])):
|
||||
if message.get('role') != 'user':
|
||||
continue
|
||||
|
||||
content = get_content_from_message(message)
|
||||
if isinstance(content, str) and content.strip():
|
||||
user_messages.append(content.strip())
|
||||
|
||||
if len(user_messages) >= 7:
|
||||
break
|
||||
|
||||
query = '\n\n'.join(reversed(user_messages))[-4000:]
|
||||
if not query:
|
||||
return form_data
|
||||
|
||||
try:
|
||||
from open_webui.routers.memories import QueryMemoryForm, query_memory
|
||||
|
||||
results = await query_memory(request, QueryMemoryForm(content=query, k=6), user)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
return form_data
|
||||
|
||||
sections = {'user': [], 'context': []}
|
||||
if results and hasattr(results, 'documents') and results.documents:
|
||||
for doc_idx, doc in enumerate(results.documents[0]):
|
||||
if not doc:
|
||||
continue
|
||||
|
||||
metadata = {}
|
||||
if results.metadatas and results.metadatas[0] and len(results.metadatas[0]) > doc_idx:
|
||||
metadata = results.metadatas[0][doc_idx] or {}
|
||||
|
||||
sections[Memories.normalize_memory_type(metadata.get('type'))].append(str(doc))
|
||||
|
||||
parts = []
|
||||
if sections['user']:
|
||||
parts.append('User:\n' + '\n'.join(f'- {memory}' for memory in sections['user']))
|
||||
if sections['context']:
|
||||
parts.append('Context:\n' + '\n'.join(f'- {memory}' for memory in sections['context']))
|
||||
if not parts:
|
||||
return form_data
|
||||
|
||||
limit = await Config.get('memories.context_char_limit', 2000)
|
||||
try:
|
||||
limit = max(250, int(limit))
|
||||
except Exception:
|
||||
limit = 2000
|
||||
|
||||
messages = form_data['messages']
|
||||
if messages and messages[0].get('role') == 'system':
|
||||
content = messages[0].get('content', '')
|
||||
if isinstance(content, str) and MEMORY_CONTEXT_OPEN in content:
|
||||
start = content.find(MEMORY_CONTEXT_OPEN)
|
||||
end = content.find(MEMORY_CONTEXT_CLOSE, start)
|
||||
if end != -1:
|
||||
messages[0]['content'] = (content[:start] + content[end + len(MEMORY_CONTEXT_CLOSE) :]).strip()
|
||||
|
||||
memory_context = f'{MEMORY_CONTEXT_OPEN}\n' + '\n\n'.join(parts)[:limit] + f'\n{MEMORY_CONTEXT_CLOSE}'
|
||||
form_data['messages'] = add_or_update_system_message(memory_context, messages, append=True)
|
||||
return form_data
|
||||
|
||||
|
||||
async def review_memory_after_turn(
|
||||
*,
|
||||
request,
|
||||
user,
|
||||
model: dict | None,
|
||||
metadata: dict,
|
||||
form_data: dict,
|
||||
assistant_message: dict,
|
||||
messages: list[dict],
|
||||
) -> None:
|
||||
if not model_allows_memory(model):
|
||||
return
|
||||
|
||||
features = metadata.get('features') or {}
|
||||
if not features.get('memory'):
|
||||
return
|
||||
|
||||
assistant_content = assistant_message.get('content', '')
|
||||
if not isinstance(assistant_content, str) or not assistant_content.strip():
|
||||
return
|
||||
|
||||
config = await Config.get_many(
|
||||
'memories.background_review.enable',
|
||||
'memories.review_interval_turns',
|
||||
)
|
||||
if not config.get('memories.background_review.enable'):
|
||||
return
|
||||
|
||||
try:
|
||||
interval = max(1, int(config.get('memories.review_interval_turns', 10)))
|
||||
except Exception:
|
||||
interval = 10
|
||||
|
||||
user_turns = len([message for message in messages if message.get('role') == 'user'])
|
||||
if user_turns == 0 or user_turns % interval != 0:
|
||||
return
|
||||
|
||||
task = asyncio.create_task(
|
||||
_review_memory(
|
||||
request=request,
|
||||
user=user,
|
||||
model=model,
|
||||
metadata=metadata,
|
||||
form_data=form_data,
|
||||
assistant_message=assistant_message,
|
||||
messages=messages,
|
||||
)
|
||||
)
|
||||
|
||||
def log_failure(done_task):
|
||||
try:
|
||||
done_task.result()
|
||||
except Exception as e:
|
||||
log.debug(f'Memory review failed: {e}')
|
||||
|
||||
task.add_done_callback(log_failure)
|
||||
|
||||
|
||||
async def _review_memory(
|
||||
*,
|
||||
request,
|
||||
user,
|
||||
model: dict | None,
|
||||
metadata: dict,
|
||||
form_data: dict,
|
||||
assistant_message: dict,
|
||||
messages: list[dict],
|
||||
) -> None:
|
||||
existing_memories = await Memories.get_memories_by_user_id(user.id)
|
||||
existing_lines = [
|
||||
f'- id={memory.id} type={memory.type} content={memory.content}'
|
||||
for memory in (existing_memories or [])[:80]
|
||||
]
|
||||
|
||||
assistant_content = assistant_message.get('content', '')
|
||||
if not isinstance(assistant_content, str):
|
||||
assistant_content = get_content_from_message(assistant_message)
|
||||
|
||||
transcript_lines = []
|
||||
for message in messages[-16:]:
|
||||
role = message.get('role', '')
|
||||
content = message.get('content', '')
|
||||
if not isinstance(content, str):
|
||||
content = get_content_from_message(message)
|
||||
content = content.strip()
|
||||
if role not in {'user', 'assistant'} or not content:
|
||||
continue
|
||||
if len(content) > 1600:
|
||||
content = f'{content[:1000]}\n...(truncated)...\n{content[-400:]}'
|
||||
transcript_lines.append(f'{role}: {content}')
|
||||
|
||||
if assistant_content.strip():
|
||||
assistant_final = assistant_content.strip()
|
||||
if len(assistant_final) > 1600:
|
||||
assistant_final = f'{assistant_final[:1000]}\n...(truncated)...\n{assistant_final[-400:]}'
|
||||
transcript_lines.append(f'assistant_final: {assistant_final}')
|
||||
|
||||
model_id = model.get('id') if isinstance(model, dict) else form_data.get('model')
|
||||
operations = await _generate_memory_operations(
|
||||
request=request,
|
||||
user=user,
|
||||
model_id=model_id,
|
||||
metadata=metadata,
|
||||
existing_text='\n'.join(existing_lines) if existing_lines else '(none)',
|
||||
transcript='\n\n'.join(transcript_lines),
|
||||
)
|
||||
if operations:
|
||||
from open_webui.routers.memories import UpdateMemoriesForm, update_memories
|
||||
|
||||
await update_memories(request, UpdateMemoriesForm(operations=operations), user)
|
||||
|
||||
|
||||
async def _generate_memory_operations(
|
||||
*,
|
||||
request,
|
||||
user,
|
||||
model_id: str,
|
||||
metadata: dict,
|
||||
existing_text: str,
|
||||
transcript: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
from open_webui.utils.chat import generate_chat_completion
|
||||
|
||||
review_prompt = f"""Review the completed conversation turn and decide whether long-term memory should change.
|
||||
|
||||
Memory types:
|
||||
- user: durable facts, preferences, or instructions about the user.
|
||||
- context: other durable context that may help future chats for this user account.
|
||||
|
||||
Rules:
|
||||
- Save only information likely to matter in future chats.
|
||||
- Do not save secrets, credentials, transient task steps, or unsupported guesses.
|
||||
- Prefer replace/remove over duplicate add when an existing memory should change.
|
||||
- Do not invent type, status, trait, score, importance, or stability schemas.
|
||||
- Return only JSON in this shape:
|
||||
{{"operations":[
|
||||
{{"action":"add","type":"user|context","content":"..."}},
|
||||
{{"action":"replace","id":"...","type":"user|context","content":"..."}},
|
||||
{{"action":"remove","id":"..."}}
|
||||
]}}
|
||||
- Use an empty operations array if nothing should be remembered.
|
||||
|
||||
Existing memories:
|
||||
{existing_text}
|
||||
|
||||
Conversation:
|
||||
{transcript}
|
||||
"""
|
||||
|
||||
response = await generate_chat_completion(
|
||||
request,
|
||||
form_data={
|
||||
'model': model_id,
|
||||
'messages': [
|
||||
{
|
||||
'role': 'system',
|
||||
'content': "You are Open WebUI's private memory reviewer. Return only valid JSON.",
|
||||
},
|
||||
{'role': 'user', 'content': review_prompt},
|
||||
],
|
||||
'stream': False,
|
||||
'metadata': {
|
||||
'task': 'memory_review',
|
||||
'chat_id': metadata.get('chat_id'),
|
||||
'message_id': metadata.get('message_id'),
|
||||
},
|
||||
},
|
||||
user=user,
|
||||
)
|
||||
|
||||
if not isinstance(response, dict) or not response.get('choices'):
|
||||
return []
|
||||
|
||||
response_message = response.get('choices', [{}])[0].get('message', {})
|
||||
content = response_message.get('content') or response_message.get('reasoning_content') or ''
|
||||
start = content.find('{')
|
||||
end = content.rfind('}')
|
||||
if start == -1 or end == -1 or end < start:
|
||||
return []
|
||||
|
||||
try:
|
||||
parsed = json.loads(content[start : end + 1])
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
operations = parsed.get('operations') if isinstance(parsed, dict) else None
|
||||
return operations if isinstance(operations, list) else []
|
||||
|
|
@ -53,7 +53,6 @@ from open_webui.routers.images import (
|
|||
image_edits,
|
||||
image_generations,
|
||||
)
|
||||
from open_webui.routers.memories import QueryMemoryForm, query_memory
|
||||
from open_webui.routers.pipelines import (
|
||||
process_pipeline_inlet_filter,
|
||||
process_pipeline_outlet_filter,
|
||||
|
|
@ -90,6 +89,7 @@ from open_webui.utils.filter import (
|
|||
)
|
||||
|
||||
from open_webui.utils.mcp.client import MCPClient
|
||||
from open_webui.utils.memory import add_memory_context, review_memory_after_turn
|
||||
from open_webui.utils.misc import (
|
||||
add_or_update_system_message,
|
||||
add_or_update_user_message,
|
||||
|
|
@ -147,6 +147,7 @@ DEFAULT_REASONING_TAGS = [
|
|||
('<|begin_of_thought|>', '<|end_of_thought|>'),
|
||||
('◁think▷', '◁/think▷'),
|
||||
]
|
||||
|
||||
DEFAULT_SOLUTION_TAGS = [('<|begin_of_solution|>', '<|end_of_solution|>')]
|
||||
DEFAULT_CODE_INTERPRETER_TAGS = [('<code_interpreter>', '</code_interpreter>')]
|
||||
|
||||
|
|
@ -1457,42 +1458,6 @@ async def chat_completion_tools_handler(
|
|||
|
||||
return body, {'sources': sources}
|
||||
|
||||
|
||||
async def chat_memory_handler(request: Request, form_data: dict, extra_params: dict, user):
|
||||
try:
|
||||
results = await query_memory(
|
||||
request,
|
||||
QueryMemoryForm(
|
||||
**{
|
||||
'content': get_last_user_message(form_data['messages']) or '',
|
||||
'k': 3,
|
||||
}
|
||||
),
|
||||
user,
|
||||
)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
results = None
|
||||
|
||||
user_context = ''
|
||||
if results and hasattr(results, 'documents'):
|
||||
if results.documents and len(results.documents) > 0:
|
||||
for doc_idx, doc in enumerate(results.documents[0]):
|
||||
created_at_date = 'Unknown Date'
|
||||
|
||||
if results.metadatas[0][doc_idx].get('created_at'):
|
||||
created_at_timestamp = results.metadatas[0][doc_idx]['created_at']
|
||||
created_at_date = time.strftime('%Y-%m-%d', time.localtime(created_at_timestamp))
|
||||
|
||||
user_context += f'{doc_idx + 1}. [{created_at_date}] {doc}\n'
|
||||
|
||||
form_data['messages'] = add_or_update_system_message(
|
||||
f'User Context:\n{user_context}\n', form_data['messages'], append=True
|
||||
)
|
||||
|
||||
return form_data
|
||||
|
||||
|
||||
async def chat_web_search_handler(request: Request, form_data: dict, extra_params: dict, user):
|
||||
event_emitter = extra_params['__event_emitter__']
|
||||
await event_emitter(
|
||||
|
|
@ -2649,8 +2614,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
|||
)
|
||||
|
||||
if 'memory' in features and features['memory']:
|
||||
# Skip forced memory injection when native FC is enabled - model can use memory tools
|
||||
form_data = await chat_memory_handler(request, form_data, user)
|
||||
form_data = await add_memory_context(request, form_data, user, model)
|
||||
|
||||
if 'web_search' in features and features['web_search']:
|
||||
# Skip forced RAG web search when native FC is enabled - model can use web_search tool
|
||||
|
|
@ -2787,6 +2751,7 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
|||
'tool_ids': tool_ids,
|
||||
'terminal_id': terminal_id,
|
||||
'files': files,
|
||||
'features': features,
|
||||
}
|
||||
form_data['metadata'] = metadata
|
||||
|
||||
|
|
|
|||
|
|
@ -82,6 +82,7 @@ from open_webui.tools.builtin import (
|
|||
toggle_automation,
|
||||
update_automation,
|
||||
update_calendar_event,
|
||||
update_memory,
|
||||
update_task,
|
||||
view_channel_message,
|
||||
view_channel_thread,
|
||||
|
|
@ -527,19 +528,21 @@ async def get_builtin_tools(
|
|||
if is_builtin_tool_enabled('chats'):
|
||||
builtin_functions.extend([search_chats, view_chat])
|
||||
|
||||
# Add memory tools if builtin category enabled AND enabled for this chat
|
||||
# Add memory tools when memory is enabled and the model allows this builtin category.
|
||||
if (
|
||||
is_builtin_tool_enabled('memory')
|
||||
and (features.get('memory') or get_model_capability('memory', False))
|
||||
and features.get('memory')
|
||||
and get_model_capability('memory')
|
||||
and await has_user_permission('memories')
|
||||
):
|
||||
builtin_functions.extend(
|
||||
[
|
||||
search_memories,
|
||||
list_memories,
|
||||
update_memory,
|
||||
add_memory,
|
||||
replace_memory_content,
|
||||
delete_memory,
|
||||
list_memories,
|
||||
]
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue