Skip to content

Commit b7ab458

Browse files
author
Ilyas Gasanov
committed
[DOP-19781] Add full-text search to users
1 parent e6c065f commit b7ab458

8 files changed

Lines changed: 213 additions & 76 deletions

File tree

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
Add full-text search for **users**

syncmaster/backend/api/v1/users.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,11 +22,17 @@ async def get_users(
2222
page_size: int = Query(gt=0, le=200, default=20),
2323
current_user: User = Depends(get_user(is_active=True)),
2424
unit_of_work: UnitOfWork = Depends(UnitOfWorkMarker),
25+
search_query: str | None = Query(
26+
None,
27+
title="Search Query",
28+
description="full-text search for users",
29+
),
2530
) -> UserPageSchema:
2631
pagination = await unit_of_work.user.paginate(
2732
page=page,
2833
page_size=page_size,
2934
is_superuser=current_user.is_superuser,
35+
search_query=search_query,
3036
)
3137
return UserPageSchema.from_pagination(pagination)
3238

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
# SPDX-FileCopyrightText: 2023-2024 MTS PJSC
2+
# SPDX-License-Identifier: Apache-2.0
3+
"""add search_vector to user table
4+
5+
Revision ID: 1df3778daf4f
6+
Revises: b9f5c4315bb2
7+
Create Date: 2024-10-07 06:59:12.667067
8+
9+
"""
10+
import sqlalchemy as sa
11+
from alembic import op
12+
from sqlalchemy.dialects import postgresql
13+
14+
# revision identifiers, used by Alembic.
15+
revision = "1df3778daf4f"
16+
down_revision = "b9f5c4315bb2"
17+
branch_labels = None
18+
depends_on = None
19+
20+
21+
def upgrade() -> None:
22+
# ### commands auto generated by Alembic - please adjust! ###
23+
op.add_column(
24+
"user",
25+
sa.Column(
26+
"search_vector",
27+
postgresql.TSVECTOR(),
28+
sa.Computed("to_tsvector('english'::regconfig, username)", persisted=True),
29+
nullable=False,
30+
),
31+
)
32+
# ### end Alembic commands ###
33+
34+
35+
def downgrade() -> None:
36+
# ### commands auto generated by Alembic - please adjust! ###
37+
op.drop_column("user", "search_vector")
38+
# ### end Alembic commands ###

syncmaster/db/models.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,16 @@ class User(Base, TimestampMixin, DeletableMixin):
3737
username: Mapped[str] = mapped_column(String(256), nullable=False, unique=True, index=True)
3838
is_superuser: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
3939
is_active: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
40+
search_vector: Mapped[str] = mapped_column(
41+
TSVECTOR,
42+
Computed(
43+
"to_tsvector('english'::regconfig, username)",
44+
persisted=True,
45+
),
46+
nullable=False,
47+
deferred=True,
48+
doc="Full-text search vector",
49+
)
4050

4151
def __repr__(self) -> str:
4252
return f"User(username={self.username}, is_superuser={self.is_superuser}, is_active={self.is_active})"

syncmaster/db/repositories/user.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
# SPDX-License-Identifier: Apache-2.0
33
from typing import Any, NoReturn
44

5-
from sqlalchemy import ScalarResult, insert, select
5+
from sqlalchemy import ScalarResult, func, insert, select
66
from sqlalchemy.exc import DBAPIError, IntegrityError, NoResultFound
77
from sqlalchemy.ext.asyncio import AsyncSession
88

@@ -17,11 +17,24 @@ class UserRepository(Repository[User]):
1717
def __init__(self, session: AsyncSession) -> None:
1818
super().__init__(model=User, session=session)
1919

20-
async def paginate(self, page: int, page_size: int, is_superuser: bool) -> Pagination:
20+
async def paginate(
21+
self,
22+
page: int,
23+
page_size: int,
24+
is_superuser: bool,
25+
search_query: str | None = None,
26+
) -> Pagination:
2127
stmt = select(User).where(User.is_deleted.is_(False))
2228

29+
if search_query:
30+
ts_query = func.plainto_tsquery("english", search_query)
31+
stmt = stmt.where(User.search_vector.op("@@")(ts_query))
32+
stmt = stmt.add_columns(func.ts_rank(User.search_vector, ts_query).label("rank"))
33+
stmt = stmt.order_by(func.ts_rank(User.search_vector, ts_query).desc())
34+
2335
if not is_superuser:
2436
stmt = stmt.where(User.is_active.is_(True))
37+
2538
return await self._paginate_scalar_result(query=stmt.order_by(User.username), page=page, page_size=page_size)
2639

2740
async def read_by_id(self, user_id: int, **kwargs: Any) -> User:

tests/test_unit/tests_users/__init__.py

Whitespace-only changes.

tests/test_unit/test_users.py renamed to tests/test_unit/tests_users/test_read_user.py

Lines changed: 23 additions & 74 deletions
Original file line numberDiff line numberDiff line change
@@ -6,13 +6,13 @@
66
pytestmark = [pytest.mark.asyncio, pytest.mark.backend]
77

88

9-
async def test_get_users(
9+
async def test_get_user(
1010
client: AsyncClient,
1111
simple_user: MockUser,
1212
inactive_user: MockUser,
1313
deleted_user: MockUser,
1414
):
15-
response = await client.get("v1/users")
15+
response = await client.get(f"/v1/users/{simple_user.id}")
1616
assert response.status_code == 401
1717
assert response.json() == {
1818
"error": {
@@ -22,72 +22,35 @@ async def test_get_users(
2222
},
2323
}
2424

25+
# check from simple user
2526
response = await client.get(
26-
"v1/users",
27+
f"v1/users/{simple_user.id}",
2728
headers={"Authorization": f"Bearer {simple_user.token}"},
2829
)
2930
assert response.status_code == 200
30-
result = response.json()
31-
assert result.keys() == {"items", "meta"}
32-
assert result["items"][0].keys() == {"id", "is_superuser", "username"}
33-
assert result["meta"] == {
34-
"page": 1,
35-
"pages": 1,
36-
"total": len(result["items"]),
37-
"page_size": 20,
38-
"has_next": False,
39-
"has_previous": False,
40-
"next_page": None,
41-
"previous_page": None,
31+
assert response.json() == {
32+
"id": simple_user.id,
33+
"is_superuser": simple_user.is_superuser,
34+
"username": simple_user.username,
4235
}
43-
for user_data in result["items"]:
44-
assert user_data["username"] != deleted_user.username
4536

37+
# check from simple user deleted user
4638
response = await client.get(
47-
"v1/users",
48-
headers={"Authorization": f"Bearer {inactive_user.token}"},
39+
f"v1/users/{deleted_user.id}",
40+
headers={"Authorization": f"Bearer {simple_user.token}"},
4941
)
50-
assert response.status_code == 403
51-
assert response.json() == {
52-
"error": {
53-
"code": "forbidden",
54-
"message": "Inactive user",
55-
"details": None,
56-
},
57-
}
58-
59-
60-
async def test_get_current_user(
61-
client: AsyncClient,
62-
simple_user: MockUser,
63-
inactive_user: MockUser,
64-
):
65-
# not authenticated user
66-
response = await client.get("/v1/users/me")
67-
assert response.status_code == 401
42+
assert response.status_code == 404
6843
assert response.json() == {
6944
"error": {
70-
"code": "unauthorized",
71-
"message": "Not authenticated",
45+
"code": "not_found",
46+
"message": "User not found",
7247
"details": None,
7348
},
7449
}
7550

76-
# active user
77-
response = await client.get(
78-
"/v1/users/me",
79-
headers={"Authorization": f"Bearer {simple_user.token}"},
80-
)
81-
assert response.status_code == 200
82-
assert response.json() == {
83-
"id": simple_user.id,
84-
"is_superuser": simple_user.is_superuser,
85-
"username": simple_user.username,
86-
}
87-
88-
# inactive user
51+
# check from inactive user
8952
response = await client.get(
90-
"/v1/users/me",
53+
f"v1/users/{simple_user.id}",
9154
headers={"Authorization": f"Bearer {inactive_user.token}"},
9255
)
9356
assert response.status_code == 403
@@ -100,13 +63,13 @@ async def test_get_current_user(
10063
}
10164

10265

103-
async def test_get_user(
66+
async def test_get_current_user(
10467
client: AsyncClient,
10568
simple_user: MockUser,
10669
inactive_user: MockUser,
107-
deleted_user: MockUser,
10870
):
109-
response = await client.get(f"/v1/users/{simple_user.id}")
71+
# not authenticated user
72+
response = await client.get("/v1/users/me")
11073
assert response.status_code == 401
11174
assert response.json() == {
11275
"error": {
@@ -116,9 +79,9 @@ async def test_get_user(
11679
},
11780
}
11881

119-
# check from simple user
82+
# active user
12083
response = await client.get(
121-
f"v1/users/{simple_user.id}",
84+
"/v1/users/me",
12285
headers={"Authorization": f"Bearer {simple_user.token}"},
12386
)
12487
assert response.status_code == 200
@@ -128,23 +91,9 @@ async def test_get_user(
12891
"username": simple_user.username,
12992
}
13093

131-
# check from simple user deleted user
132-
response = await client.get(
133-
f"v1/users/{deleted_user.id}",
134-
headers={"Authorization": f"Bearer {simple_user.token}"},
135-
)
136-
assert response.status_code == 404
137-
assert response.json() == {
138-
"error": {
139-
"code": "not_found",
140-
"message": "User not found",
141-
"details": None,
142-
},
143-
}
144-
145-
# check from inactive user
94+
# inactive user
14695
response = await client.get(
147-
f"v1/users/{simple_user.id}",
96+
"/v1/users/me",
14897
headers={"Authorization": f"Bearer {inactive_user.token}"},
14998
)
15099
assert response.status_code == 403
Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,120 @@
1+
import random
2+
import string
3+
4+
import pytest
5+
from httpx import AsyncClient
6+
7+
from tests.mocks import MockUser
8+
9+
pytestmark = [pytest.mark.asyncio, pytest.mark.backend]
10+
11+
12+
async def test_get_users(
13+
client: AsyncClient,
14+
simple_user: MockUser,
15+
inactive_user: MockUser,
16+
deleted_user: MockUser,
17+
):
18+
response = await client.get("v1/users")
19+
assert response.status_code == 401
20+
assert response.json() == {
21+
"error": {
22+
"code": "unauthorized",
23+
"message": "Not authenticated",
24+
"details": None,
25+
},
26+
}
27+
28+
response = await client.get(
29+
"v1/users",
30+
headers={"Authorization": f"Bearer {simple_user.token}"},
31+
)
32+
assert response.status_code == 200
33+
result = response.json()
34+
assert result.keys() == {"items", "meta"}
35+
assert result["items"][0].keys() == {"id", "is_superuser", "username"}
36+
assert result["meta"] == {
37+
"page": 1,
38+
"pages": 1,
39+
"total": len(result["items"]),
40+
"page_size": 20,
41+
"has_next": False,
42+
"has_previous": False,
43+
"next_page": None,
44+
"previous_page": None,
45+
}
46+
for user_data in result["items"]:
47+
assert user_data["username"] != deleted_user.username
48+
49+
response = await client.get(
50+
"v1/users",
51+
headers={"Authorization": f"Bearer {inactive_user.token}"},
52+
)
53+
assert response.status_code == 403
54+
assert response.json() == {
55+
"error": {
56+
"code": "forbidden",
57+
"message": "Inactive user",
58+
"details": None,
59+
},
60+
}
61+
62+
63+
@pytest.mark.parametrize(
64+
"search_value_extractor",
65+
[
66+
lambda user: user.username,
67+
],
68+
ids=["search_by_username"],
69+
)
70+
async def test_search_users_with_query(
71+
client: AsyncClient,
72+
superuser: MockUser,
73+
simple_user: MockUser,
74+
search_value_extractor,
75+
):
76+
user = simple_user
77+
search_query = search_value_extractor(user)
78+
79+
result = await client.get(
80+
"v1/users",
81+
headers={"Authorization": f"Bearer {superuser.token}"},
82+
params={"search_query": search_query},
83+
)
84+
85+
assert result.json() == {
86+
"meta": {
87+
"page": 1,
88+
"pages": 1,
89+
"total": 1,
90+
"page_size": 20,
91+
"has_next": False,
92+
"has_previous": False,
93+
"next_page": None,
94+
"previous_page": None,
95+
},
96+
"items": [
97+
{
98+
"id": user.id,
99+
"username": user.username,
100+
"is_superuser": user.is_superuser,
101+
},
102+
],
103+
}
104+
assert result.status_code == 200
105+
106+
107+
async def test_search_users_with_nonexistent_query(
108+
client: AsyncClient,
109+
superuser: MockUser,
110+
):
111+
random_search_query = "".join(random.choices(string.ascii_lowercase + string.digits, k=12))
112+
113+
result = await client.get(
114+
"v1/users",
115+
headers={"Authorization": f"Bearer {superuser.token}"},
116+
params={"search_query": random_search_query},
117+
)
118+
119+
assert result.status_code == 200
120+
assert result.json()["items"] == []

0 commit comments

Comments
 (0)