Skip to content

Commit 260f3c3

Browse files
committed
refac
1 parent b0487dd commit 260f3c3

7 files changed

Lines changed: 302 additions & 259 deletions

File tree

backend/open_webui/routers/memories.py

Lines changed: 102 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
import asyncio
44
import logging
5-
from typing import Literal, Optional
5+
from typing import Literal
66

77
from fastapi import APIRouter, Depends, HTTPException, Request, status
88
from open_webui.constants import ERROR_MESSAGES
@@ -14,7 +14,7 @@
1414
from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
1515
from open_webui.utils.access_control import has_permission
1616
from open_webui.utils.auth import get_verified_user
17-
from open_webui.utils.memory import clean_memory_content, validate_memory_operations
17+
from open_webui.utils.memory import clean_memory_content, clean_memory_path, memory_vector_text, validate_memory_operations
1818
from pydantic import BaseModel
1919
from sqlalchemy.ext.asyncio import AsyncSession
2020

@@ -66,22 +66,43 @@ async def get_memories(
6666
class AddMemoryForm(BaseModel):
6767
content: str
6868
type: Literal['user', 'context'] = 'context'
69+
path: str | None = None
6970

7071

7172
class MemoryUpdateModel(BaseModel):
7273
content: str | None = None
7374
type: Literal['user', 'context'] | None = None
75+
path: str | None = None
7476

7577

7678
class MemoryOperationModel(BaseModel):
77-
action: Literal['add', 'replace', 'remove']
79+
action: Literal['add', 'replace', 'remove', 'move']
7880
id: str | None = None
7981
content: str | None = None
8082
type: Literal['user', 'context'] | None = None
83+
path: str | None = None
8184

8285

8386
class UpdateMemoriesForm(BaseModel):
8487
operations: list[MemoryOperationModel]
88+
source: Literal['tool', 'background_review'] | None = None
89+
90+
91+
class SearchMemoriesForm(BaseModel):
92+
query: str | None = None
93+
type: Literal['user', 'context', 'all'] = 'all'
94+
path: str | None = None
95+
memory_id: str | None = None
96+
limit: int = 20
97+
98+
99+
def _memory_metadata(memory: MemoryModel) -> dict:
100+
return {
101+
'created_at': memory.created_at,
102+
'updated_at': memory.updated_at,
103+
'type': memory.type,
104+
'path': memory.path,
105+
}
85106

86107

87108
@router.post('/add', response_model=MemoryModel | None)
@@ -99,22 +120,25 @@ async def add_memory(
99120
await check_memories_permission(user)
100121

101122
content = clean_memory_content(form_data.content)
102-
memory = await Memories.insert_new_memory(user.id, content, memory_type=form_data.type)
123+
path = clean_memory_path(form_data.path)
124+
memory = await Memories.insert_new_memory(
125+
user.id,
126+
content,
127+
memory_type=form_data.type,
128+
path=path,
129+
meta={'created_by': 'manual'},
130+
)
103131

104-
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
132+
vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
105133

106134
await ASYNC_VECTOR_DB_CLIENT.upsert(
107135
collection_name=f'user-memory-{user.id}',
108136
items=[
109137
{
110138
'id': memory.id,
111-
'text': memory.content,
139+
'text': memory_vector_text(memory.content, memory.path),
112140
'vector': vector,
113-
'metadata': {
114-
'created_at': memory.created_at,
115-
'updated_at': memory.updated_at,
116-
'type': memory.type,
117-
},
141+
'metadata': _memory_metadata(memory),
118142
}
119143
],
120144
)
@@ -124,7 +148,7 @@ async def add_memory(
124148
EVENTS.MEMORY_CREATED,
125149
actor=user,
126150
subject_id=memory.id,
127-
data={'content_preview': memory.content[:300], 'type': memory.type},
151+
data={'content_preview': memory.content[:300], 'type': memory.type, 'path': memory.path},
128152
)
129153
return memory
130154

@@ -138,6 +162,16 @@ async def update_memories(
138162
await check_memories_permission(user)
139163

140164
operations = validate_memory_operations(form_data)
165+
metadata = getattr(request.state, 'metadata', {}) or {}
166+
source = form_data.source or 'tool'
167+
for operation in operations:
168+
if operation.get('action') in {'add', 'replace', 'move'}:
169+
operation['meta'] = {
170+
'created_by': source,
171+
'chat_id': metadata.get('chat_id'),
172+
'message_id': metadata.get('message_id'),
173+
'model': metadata.get('model'),
174+
}
141175

142176
try:
143177
results = await Memories.apply_memory_operations(user.id, operations)
@@ -153,17 +187,16 @@ async def update_memories(
153187
if isinstance(memory, MemoryModel):
154188
result = {**result, 'memory': memory.model_dump()}
155189
if result.get('status') in {'created', 'updated'}:
156-
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
190+
vector = await request.app.state.EMBEDDING_FUNCTION(
191+
memory_vector_text(memory.content, memory.path),
192+
user=user,
193+
)
157194
upsert_items.append(
158195
{
159196
'id': memory.id,
160-
'text': memory.content,
197+
'text': memory_vector_text(memory.content, memory.path),
161198
'vector': vector,
162-
'metadata': {
163-
'created_at': memory.created_at,
164-
'updated_at': memory.updated_at,
165-
'type': memory.type,
166-
},
199+
'metadata': _memory_metadata(memory),
167200
}
168201
)
169202
if result.get('status') == 'deleted' and result.get('id'):
@@ -198,6 +231,7 @@ async def update_memories(
198231
data={
199232
'content_preview': (memory.get('content') or '')[:300],
200233
'type': memory.get('type'),
234+
'path': memory.get('path'),
201235
'operation': result.get('action'),
202236
},
203237
)
@@ -274,6 +308,33 @@ async def query_memory(
274308
return results
275309

276310

311+
@router.post('/search', response_model=list[MemoryModel])
312+
async def search_memories(
313+
form_data: SearchMemoriesForm,
314+
user=Depends(get_verified_user),
315+
):
316+
await check_memories_permission(user)
317+
318+
memories = await Memories.get_memories_by_user_id(user.id)
319+
if form_data.memory_id:
320+
memories = [memory for memory in memories if memory.id == form_data.memory_id]
321+
if form_data.type != 'all':
322+
memories = [memory for memory in memories if memory.type == form_data.type]
323+
path = clean_memory_path(form_data.path)
324+
if path:
325+
memories = [memory for memory in memories if (memory.path or '').startswith(path)]
326+
query = (form_data.query or '').strip().lower()
327+
if query:
328+
memories = [
329+
memory
330+
for memory in memories
331+
if query in memory.content.lower() or query in (memory.path or '').lower()
332+
]
333+
334+
limit = max(1, min(form_data.limit or 20, 100))
335+
return sorted(memories, key=lambda memory: memory.updated_at, reverse=True)[:limit]
336+
337+
277338
############################
278339
# ResetMemoryFromVectorDB
279340
############################
@@ -298,21 +359,20 @@ async def reset_memory_from_vector_db(
298359

299360
# Generate vectors in parallel
300361
vectors = await asyncio.gather(
301-
*[request.app.state.EMBEDDING_FUNCTION(memory.content, user=user) for memory in memories]
362+
*[
363+
request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
364+
for memory in memories
365+
]
302366
)
303367

304368
await ASYNC_VECTOR_DB_CLIENT.upsert(
305369
collection_name=f'user-memory-{user.id}',
306370
items=[
307371
{
308372
'id': memory.id,
309-
'text': memory.content,
373+
'text': memory_vector_text(memory.content, memory.path),
310374
'vector': vectors[idx],
311-
'metadata': {
312-
'created_at': memory.created_at,
313-
'updated_at': memory.updated_at,
314-
'type': memory.type,
315-
},
375+
'metadata': _memory_metadata(memory),
316376
}
317377
for idx, memory in enumerate(memories)
318378
],
@@ -378,27 +438,32 @@ async def update_memory_by_id(
378438
await check_memories_permission(user)
379439

380440
content = clean_memory_content(form_data.content) if form_data.content is not None else None
381-
if content is None and form_data.type is None:
441+
path = clean_memory_path(form_data.path)
442+
if content is None and form_data.type is None and form_data.path is None:
382443
raise HTTPException(status_code=400, detail='No memory update provided')
383-
memory = await Memories.update_memory_by_id_and_user_id(memory_id, user.id, content, memory_type=form_data.type)
444+
memory = await Memories.update_memory_by_id_and_user_id(
445+
memory_id,
446+
user.id,
447+
content,
448+
memory_type=form_data.type,
449+
path=path,
450+
update_path=form_data.path is not None,
451+
meta={'created_by': 'manual'},
452+
)
384453
if memory is None:
385454
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
386455

387-
if form_data.content is not None:
388-
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
456+
if form_data.content is not None or form_data.path is not None:
457+
vector = await request.app.state.EMBEDDING_FUNCTION(memory_vector_text(memory.content, memory.path), user=user)
389458

390459
await ASYNC_VECTOR_DB_CLIENT.upsert(
391460
collection_name=f'user-memory-{user.id}',
392461
items=[
393462
{
394463
'id': memory.id,
395-
'text': memory.content,
464+
'text': memory_vector_text(memory.content, memory.path),
396465
'vector': vector,
397-
'metadata': {
398-
'created_at': memory.created_at,
399-
'updated_at': memory.updated_at,
400-
'type': memory.type,
401-
},
466+
'metadata': _memory_metadata(memory),
402467
}
403468
],
404469
)
@@ -408,7 +473,7 @@ async def update_memory_by_id(
408473
EVENTS.MEMORY_UPDATED,
409474
actor=user,
410475
subject_id=memory.id,
411-
data={'content_preview': memory.content[:300], 'type': memory.type},
476+
data={'content_preview': memory.content[:300], 'type': memory.type, 'path': memory.path},
412477
)
413478
return memory
414479

0 commit comments

Comments
 (0)