refac
This commit is contained in:
@@ -23,7 +23,7 @@ router = APIRouter()
|
||||
############################
|
||||
|
||||
|
||||
@router.get("/", response_model=list[MemoryModel])
|
||||
@router.get('/', response_model=list[MemoryModel])
|
||||
async def get_memories(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -35,9 +35,7 @@ async def get_memories(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -59,7 +57,7 @@ class MemoryUpdateModel(BaseModel):
|
||||
content: Optional[str] = None
|
||||
|
||||
|
||||
@router.post("/add", response_model=Optional[MemoryModel])
|
||||
@router.post('/add', response_model=Optional[MemoryModel])
|
||||
async def add_memory(
|
||||
request: Request,
|
||||
form_data: AddMemoryForm,
|
||||
@@ -75,9 +73,7 @@ async def add_memory(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -88,13 +84,13 @@ async def add_memory(
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f"user-memory-{user.id}",
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
"id": memory.id,
|
||||
"text": memory.content,
|
||||
"vector": vector,
|
||||
"metadata": {"created_at": memory.created_at},
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'vector': vector,
|
||||
'metadata': {'created_at': memory.created_at},
|
||||
}
|
||||
],
|
||||
)
|
||||
@@ -112,7 +108,7 @@ class QueryMemoryForm(BaseModel):
|
||||
k: Optional[int] = 1
|
||||
|
||||
|
||||
@router.post("/query")
|
||||
@router.post('/query')
|
||||
async def query_memory(
|
||||
request: Request,
|
||||
form_data: QueryMemoryForm,
|
||||
@@ -128,9 +124,7 @@ async def query_memory(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -138,12 +132,12 @@ async def query_memory(
|
||||
|
||||
memories = Memories.get_memories_by_user_id(user.id)
|
||||
if not memories:
|
||||
raise HTTPException(status_code=404, detail="No memories found for user")
|
||||
raise HTTPException(status_code=404, detail='No memories found for user')
|
||||
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(form_data.content, user=user)
|
||||
|
||||
results = VECTOR_DB_CLIENT.search(
|
||||
collection_name=f"user-memory-{user.id}",
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
vectors=[vector],
|
||||
limit=form_data.k,
|
||||
)
|
||||
@@ -154,7 +148,7 @@ async def query_memory(
|
||||
############################
|
||||
# ResetMemoryFromVectorDB
|
||||
############################
|
||||
@router.post("/reset", response_model=bool)
|
||||
@router.post('/reset', response_model=bool)
|
||||
async def reset_memory_from_vector_db(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -173,36 +167,31 @@ async def reset_memory_from_vector_db(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
VECTOR_DB_CLIENT.delete_collection(f"user-memory-{user.id}")
|
||||
VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
|
||||
memories = Memories.get_memories_by_user_id(user.id)
|
||||
|
||||
# Generate vectors in parallel
|
||||
vectors = await asyncio.gather(
|
||||
*[
|
||||
request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
for memory in memories
|
||||
]
|
||||
*[request.app.state.EMBEDDING_FUNCTION(memory.content, user=user) for memory in memories]
|
||||
)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f"user-memory-{user.id}",
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
"id": memory.id,
|
||||
"text": memory.content,
|
||||
"vector": vectors[idx],
|
||||
"metadata": {
|
||||
"created_at": memory.created_at,
|
||||
"updated_at": memory.updated_at,
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'vector': vectors[idx],
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
},
|
||||
}
|
||||
for idx, memory in enumerate(memories)
|
||||
@@ -217,7 +206,7 @@ async def reset_memory_from_vector_db(
|
||||
############################
|
||||
|
||||
|
||||
@router.delete("/delete/user", response_model=bool)
|
||||
@router.delete('/delete/user', response_model=bool)
|
||||
async def delete_memory_by_user_id(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
@@ -229,9 +218,7 @@ async def delete_memory_by_user_id(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -241,7 +228,7 @@ async def delete_memory_by_user_id(
|
||||
|
||||
if result:
|
||||
try:
|
||||
VECTOR_DB_CLIENT.delete_collection(f"user-memory-{user.id}")
|
||||
VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}')
|
||||
except Exception as e:
|
||||
log.error(e)
|
||||
return True
|
||||
@@ -254,7 +241,7 @@ async def delete_memory_by_user_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.post("/{memory_id}/update", response_model=Optional[MemoryModel])
|
||||
@router.post('/{memory_id}/update', response_model=Optional[MemoryModel])
|
||||
async def update_memory_by_id(
|
||||
memory_id: str,
|
||||
request: Request,
|
||||
@@ -271,33 +258,29 @@ async def update_memory_by_id(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
memory = Memories.update_memory_by_id_and_user_id(
|
||||
memory_id, user.id, form_data.content
|
||||
)
|
||||
memory = Memories.update_memory_by_id_and_user_id(memory_id, user.id, form_data.content)
|
||||
if memory is None:
|
||||
raise HTTPException(status_code=404, detail="Memory not found")
|
||||
raise HTTPException(status_code=404, detail='Memory not found')
|
||||
|
||||
if form_data.content is not None:
|
||||
vector = await request.app.state.EMBEDDING_FUNCTION(memory.content, user=user)
|
||||
|
||||
VECTOR_DB_CLIENT.upsert(
|
||||
collection_name=f"user-memory-{user.id}",
|
||||
collection_name=f'user-memory-{user.id}',
|
||||
items=[
|
||||
{
|
||||
"id": memory.id,
|
||||
"text": memory.content,
|
||||
"vector": vector,
|
||||
"metadata": {
|
||||
"created_at": memory.created_at,
|
||||
"updated_at": memory.updated_at,
|
||||
'id': memory.id,
|
||||
'text': memory.content,
|
||||
'vector': vector,
|
||||
'metadata': {
|
||||
'created_at': memory.created_at,
|
||||
'updated_at': memory.updated_at,
|
||||
},
|
||||
}
|
||||
],
|
||||
@@ -311,7 +294,7 @@ async def update_memory_by_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.delete("/{memory_id}", response_model=bool)
|
||||
@router.delete('/{memory_id}', response_model=bool)
|
||||
async def delete_memory_by_id(
|
||||
memory_id: str,
|
||||
request: Request,
|
||||
@@ -324,9 +307,7 @@ async def delete_memory_by_id(
|
||||
detail=ERROR_MESSAGES.NOT_FOUND,
|
||||
)
|
||||
|
||||
if not has_permission(
|
||||
user.id, "features.memories", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
if not has_permission(user.id, 'features.memories', request.app.state.config.USER_PERMISSIONS):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
@@ -335,9 +316,7 @@ async def delete_memory_by_id(
|
||||
result = Memories.delete_memory_by_id_and_user_id(memory_id, user.id, db=db)
|
||||
|
||||
if result:
|
||||
VECTOR_DB_CLIENT.delete(
|
||||
collection_name=f"user-memory-{user.id}", ids=[memory_id]
|
||||
)
|
||||
VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=[memory_id])
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user