From 8c632f9ef9a53b6c95726e0e6351950fd87746f9 Mon Sep 17 00:00:00 2001 From: Jordi Smit Date: Mon, 23 Jan 2023 22:46:59 +0100 Subject: [PATCH] add lang filter option to message endpoints (#902) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * add lang filter option to message endpoints * add lang filter option to frontend_users endpoints * move lang param before api_client Co-authored-by: Andreas Köpf --- backend/oasst_backend/api/v1/frontend_users.py | 4 ++++ backend/oasst_backend/api/v1/messages.py | 4 ++++ backend/oasst_backend/api/v1/users.py | 4 ++++ backend/oasst_backend/prompt_repository.py | 4 ++++ 4 files changed, 16 insertions(+) diff --git a/backend/oasst_backend/api/v1/frontend_users.py b/backend/oasst_backend/api/v1/frontend_users.py index a01e009a..5ea7b26c 100644 --- a/backend/oasst_backend/api/v1/frontend_users.py +++ b/backend/oasst_backend/api/v1/frontend_users.py @@ -70,6 +70,7 @@ def query_frontend_user_messages( only_roots: bool = False, desc: bool = True, include_deleted: bool = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -87,6 +88,7 @@ def query_frontend_user_messages( lte_created_date=end_date, only_roots=only_roots, deleted=None if include_deleted else False, + lang=lang, ) return utils.prepare_message_list(messages) @@ -101,6 +103,7 @@ def query_frontend_user_messages_cursor( include_deleted: Optional[bool] = False, max_count: Optional[int] = Query(10, gt=0, le=1000), desc: Optional[bool] = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -113,6 +116,7 @@ def query_frontend_user_messages_cursor( include_deleted=include_deleted, max_count=max_count, desc=desc, + lang=lang, api_client=api_client, db=db, ) diff --git a/backend/oasst_backend/api/v1/messages.py b/backend/oasst_backend/api/v1/messages.py index d3d5e1c3..06dd3fe1 100644 --- a/backend/oasst_backend/api/v1/messages.py +++ b/backend/oasst_backend/api/v1/messages.py @@ -26,6 +26,7 @@ def query_messages( only_roots: Optional[bool] = False, desc: Optional[bool] = True, allow_deleted: Optional[bool] = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -43,6 +44,7 @@ def query_messages( lte_created_date=end_date, only_roots=only_roots, deleted=None if allow_deleted else False, + lang=lang, ) return utils.prepare_message_list(messages) @@ -60,6 +62,7 @@ def get_messages_cursor( include_deleted: Optional[bool] = False, max_count: Optional[int] = Query(10, gt=0, le=1000), desc: Optional[bool] = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -103,6 +106,7 @@ def get_messages_cursor( deleted=None if include_deleted else False, desc=query_desc, limit=qry_max_count, + lang=lang, ) num_rows = len(items) diff --git a/backend/oasst_backend/api/v1/users.py b/backend/oasst_backend/api/v1/users.py index e4683a76..c0055339 100644 --- a/backend/oasst_backend/api/v1/users.py +++ b/backend/oasst_backend/api/v1/users.py @@ -223,6 +223,7 @@ def query_user_messages( only_roots: bool = False, desc: bool = True, include_deleted: bool = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -239,6 +240,7 @@ def query_user_messages( lte_created_date=end_date, only_roots=only_roots, deleted=None if include_deleted else False, + lang=lang, ) return utils.prepare_message_list(messages) @@ -253,6 +255,7 @@ def query_user_messages_cursor( include_deleted: Optional[bool] = False, max_count: Optional[int] = Query(10, gt=0, le=1000), desc: Optional[bool] = False, + lang: Optional[str] = None, api_client: ApiClient = Depends(deps.get_api_client), db: Session = Depends(deps.get_db), ): @@ -264,6 +267,7 @@ def query_user_messages_cursor( include_deleted=include_deleted, max_count=max_count, desc=desc, + lang=lang, api_client=api_client, db=db, ) diff --git a/backend/oasst_backend/prompt_repository.py b/backend/oasst_backend/prompt_repository.py index 0a0fa61d..a0bc2ae7 100644 --- a/backend/oasst_backend/prompt_repository.py +++ b/backend/oasst_backend/prompt_repository.py @@ -694,6 +694,7 @@ class PromptRepository: deleted: Optional[bool] = None, desc: bool = False, limit: Optional[int] = 100, + lang: Optional[str] = None, ) -> list[Message]: if not self.api_client.trusted: if not api_client_id: @@ -758,6 +759,9 @@ class PromptRepository: if limit is not None: qry = qry.limit(limit) + if lang is not None: + qry = qry.filter(Message.lang == lang) + return qry.all() def update_children_counts(self, message_tree_id: UUID):