Skip to content

Commit 826244a

Browse files
deduplication code
1 parent 4d43e78 commit 826244a

12 files changed

Lines changed: 260 additions & 154 deletions

File tree

app/api/v1/endpoints/auth.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,9 +51,12 @@ async def logout(
5151

5252

5353
@router.get("/me", response_model=MeResponse)
54-
async def get_current_user_profile(current_user: User = Depends(get_current_user)):
54+
async def get_current_user_profile(
55+
current_user: User = Depends(get_current_user),
56+
db: AsyncSession = Depends(get_db),
57+
):
5558
"""Get the current user's profile."""
56-
return auth_service.get_current_user_profile(current_user=current_user)
59+
return await auth_service.get_current_user_profile(current_user=current_user, db=db)
5760

5861

5962
@router.get("/me/sessions", response_model=list[SessionResponse])

app/api/v1/endpoints/chat.py

Lines changed: 16 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from sqlalchemy.ext.asyncio import AsyncSession
55

66
from app.api.dependencies import require_company_user
7+
from app.api.v1.endpoints.common import build_list_service_kwargs
78
from app.db.models import User
89
from app.db.session import get_db
910
from app.schemas.chat import ChatConversationResponse, ChatMessageResponse, ChatRequest, ChatResponse
@@ -54,12 +55,14 @@ async def list_conversations(
5455
current_user: User = Depends(require_company_user),
5556
):
5657
"""List persisted chat conversations for the current authenticated user."""
57-
service_kwargs: dict = {"db": db, "current_user": current_user}
58-
if limit is not None:
59-
service_kwargs["limit"] = limit
60-
if offset is not None:
61-
service_kwargs["offset"] = offset
62-
return await chat_service.list_conversations(**service_kwargs)
58+
return await chat_service.list_conversations(
59+
**build_list_service_kwargs(
60+
db=db,
61+
current_user=current_user,
62+
limit=limit,
63+
offset=offset,
64+
)
65+
)
6366

6467

6568
@router.get("/conversations/{conversation_id}/messages", response_model=list[ChatMessageResponse])
@@ -71,15 +74,12 @@ async def get_conversation_messages(
7174
current_user: User = Depends(require_company_user),
7275
):
7376
"""Get persisted messages for one conversation owned by the current user."""
74-
service_kwargs: dict = {
75-
"conversation_id": conversation_id,
76-
"db": db,
77-
"current_user": current_user,
78-
}
79-
if limit is not None:
80-
service_kwargs["limit"] = limit
81-
if offset is not None:
82-
service_kwargs["offset"] = offset
8377
return await chat_service.get_conversation_messages(
84-
**service_kwargs,
78+
**build_list_service_kwargs(
79+
db=db,
80+
current_user=current_user,
81+
limit=limit,
82+
offset=offset,
83+
conversation_id=conversation_id,
84+
),
8585
)

app/api/v1/endpoints/common.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,28 @@
1+
"""Shared endpoint-layer helpers for building service kwargs."""
2+
3+
from __future__ import annotations
4+
5+
from typing import Any
6+
7+
from app.db.models import User
8+
9+
10+
def build_list_service_kwargs(
11+
*,
12+
db: Any,
13+
current_user: User,
14+
limit: int | None = None,
15+
offset: int | None = None,
16+
**extra: Any,
17+
) -> dict[str, Any]:
18+
"""Build service kwargs while preserving optional-parameter compatibility for tests/mocks."""
19+
kwargs: dict[str, Any] = {
20+
"db": db,
21+
"current_user": current_user,
22+
**extra,
23+
}
24+
if limit is not None:
25+
kwargs["limit"] = limit
26+
if offset is not None:
27+
kwargs["offset"] = offset
28+
return kwargs

app/api/v1/endpoints/companies.py

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
from sqlalchemy.ext.asyncio import AsyncSession
55

66
from app.api.dependencies import require_super_admin
7+
from app.api.v1.endpoints.common import build_list_service_kwargs
78
from app.db.models import User
89
from app.db.session import get_db
910
from app.schemas.companies import CompanyCreate, CompanyResponse, CompanyUpdate
@@ -34,12 +35,14 @@ async def list_companies(
3435
current_user: User = Depends(require_super_admin),
3536
):
3637
"""Return all companies ordered by name (super admin only)."""
37-
service_kwargs: dict = {"db": db, "current_user": current_user}
38-
if limit is not None:
39-
service_kwargs["limit"] = limit
40-
if offset is not None:
41-
service_kwargs["offset"] = offset
42-
return await companies_service.list_companies(**service_kwargs)
38+
return await companies_service.list_companies(
39+
**build_list_service_kwargs(
40+
db=db,
41+
current_user=current_user,
42+
limit=limit,
43+
offset=offset,
44+
)
45+
)
4346

4447

4548
@router.get("/{company_id}", response_model=CompanyResponse)

app/api/v1/endpoints/documents.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from sqlalchemy.ext.asyncio import AsyncSession
33

44
from app.api.dependencies import require_admin_or_super_admin
5+
from app.api.v1.endpoints.common import build_list_service_kwargs
56
from app.db.models import User
67
from app.db.session import get_db
78
from app.schemas.documents import DocumentResponse, UploadResponse
@@ -60,16 +61,15 @@ async def list_documents(
6061
**admin**: always returns documents from their own company only.
6162
**super_admin**: returns documents for the specified ``company_id``; if omitted, returns all documents.
6263
"""
63-
service_kwargs: dict = {
64-
"company_id": company_id,
65-
"db": db,
66-
"current_user": current_user,
67-
}
68-
if limit is not None:
69-
service_kwargs["limit"] = limit
70-
if offset is not None:
71-
service_kwargs["offset"] = offset
72-
return await documents_service.list_documents(**service_kwargs)
64+
return await documents_service.list_documents(
65+
**build_list_service_kwargs(
66+
db=db,
67+
current_user=current_user,
68+
limit=limit,
69+
offset=offset,
70+
company_id=company_id,
71+
)
72+
)
7373

7474

7575
@router.get("/{document_id}", response_model=DocumentResponse, status_code=200)

app/api/v1/endpoints/users.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from sqlalchemy.ext.asyncio import AsyncSession
1818

1919
from app.api.dependencies import require_admin_or_super_admin
20+
from app.api.v1.endpoints.common import build_list_service_kwargs
2021
from app.db.models import User
2122
from app.db.session import get_db
2223
from app.schemas.users import UserCreate, UserResponse, UserUpdate
@@ -71,16 +72,15 @@ async def list_users(
7172
**super_admin**: returns users for the specified ``company_id``; if omitted, returns all users
7273
across every company (excluding other ``super_admin`` accounts).
7374
"""
74-
service_kwargs: dict = {
75-
"company_id": company_id,
76-
"db": db,
77-
"current_user": current_user,
78-
}
79-
if limit is not None:
80-
service_kwargs["limit"] = limit
81-
if offset is not None:
82-
service_kwargs["offset"] = offset
83-
return await users_service.list_users(**service_kwargs)
75+
return await users_service.list_users(
76+
**build_list_service_kwargs(
77+
db=db,
78+
current_user=current_user,
79+
limit=limit,
80+
offset=offset,
81+
company_id=company_id,
82+
)
83+
)
8484

8585

8686
@router.get("/{user_id}", response_model=UserResponse)

app/services/auth_service.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
from app.core.logger import get_logger
1313
from app.core.runtime_controls import rate_limit_exceeded, token_session_cache_delete
1414
from app.core.security import create_access_token, verify_password
15-
from app.db.models import TokenSession, User, UserRole
15+
from app.db.models import Company, TokenSession, User, UserRole
1616
from app.schemas.auth import LogoutResponse, MeResponse, TokenResponse
1717

1818
logger = get_logger(__name__)
@@ -115,15 +115,20 @@ async def logout(*, token: str, current_user: User, db: AsyncSession) -> LogoutR
115115
return LogoutResponse(message="Successfully logged out")
116116

117117

118-
def get_current_user_profile(*, current_user: User) -> MeResponse:
118+
async def get_current_user_profile(*, current_user: User, db: AsyncSession) -> MeResponse:
119119
"""Convert the authenticated ORM user into API profile response."""
120120
logger.info("Profile fetched", extra={"user_id": current_user.id, "username": current_user.username})
121+
company_name = None
122+
if current_user.company_id:
123+
company_result = await db.execute(select(Company).filter(Company.id == current_user.company_id))
124+
company = company_result.scalar_one_or_none()
125+
company_name = company.name if company else None
121126
return MeResponse(
122127
id=current_user.id,
123128
username=current_user.username,
124129
role=current_user.role,
125130
company_id=current_user.company_id,
126-
company_name=current_user.company.name if current_user.company else None,
131+
company_name=company_name,
127132
created_at=current_user.created_at,
128133
)
129134

app/services/chat_service.py

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -30,17 +30,10 @@
3030
)
3131
from app.db.models import ChatConversation, ChatMessage, User
3232
from app.schemas.chat import ChatConversationResponse, ChatMessageResponse
33+
from app.services.common import sanitize_pagination
3334

3435
logger = get_logger(__name__)
3536

36-
37-
def _sanitize_pagination(*, limit: int | None, offset: int | None) -> tuple[int, int]:
38-
safe_limit = settings.DEFAULT_LIST_LIMIT if limit is None else limit
39-
safe_limit = max(1, min(safe_limit, settings.MAX_LIST_LIMIT))
40-
safe_offset = 0 if offset is None else max(0, offset)
41-
return safe_limit, safe_offset
42-
43-
4437
def _conversation_list_cache_key(*, user_id: str, company_id: str) -> str:
4538
return f"cache:chat:conversations:{company_id}:{user_id}"
4639

@@ -270,7 +263,7 @@ async def list_conversations(
270263
detail="You do not have permission to perform this action.",
271264
)
272265

273-
safe_limit, safe_offset = _sanitize_pagination(limit=limit, offset=offset)
266+
safe_limit, safe_offset = sanitize_pagination(limit=limit, offset=offset)
274267
cache_key = f"{_conversation_list_cache_key(user_id=current_user.id, company_id=current_user.company_id)}:{safe_limit}:{safe_offset}"
275268
cached = await cache_get_json(key=cache_key)
276269
if isinstance(cached, list):
@@ -314,7 +307,7 @@ async def get_conversation_messages(
314307
"""Return persisted messages for one user-scoped conversation."""
315308
conversation = await _get_scoped_conversation(conversation_id=conversation_id, current_user=current_user, db=db)
316309

317-
safe_limit, safe_offset = _sanitize_pagination(limit=limit, offset=offset)
310+
safe_limit, safe_offset = sanitize_pagination(limit=limit, offset=offset)
318311
cache_key = f"{_conversation_messages_cache_key(conversation_id=conversation.id)}:{safe_limit}:{safe_offset}"
319312
cached = await cache_get_json(key=cache_key)
320313
if isinstance(cached, list):

app/services/common.py

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
"""Shared helpers for service-layer data access and pagination."""
2+
3+
from __future__ import annotations
4+
5+
from typing import Any
6+
7+
from fastapi import HTTPException, status
8+
from sqlalchemy import select
9+
from sqlalchemy.ext.asyncio import AsyncSession
10+
11+
from app.core.config import settings
12+
from app.core.logger import get_logger
13+
from app.db.models import Company
14+
15+
logger = get_logger(__name__)
16+
17+
18+
def sanitize_pagination(*, limit: int | None, offset: int | None) -> tuple[int, int]:
19+
"""Normalize list pagination parameters using configured defaults and limits."""
20+
safe_limit = settings.DEFAULT_LIST_LIMIT if limit is None else limit
21+
safe_limit = max(1, min(safe_limit, settings.MAX_LIST_LIMIT))
22+
safe_offset = 0 if offset is None else max(0, offset)
23+
return safe_limit, safe_offset
24+
25+
26+
async def get_by_id_or_404(
27+
*,
28+
db: AsyncSession,
29+
model: type[Any],
30+
entity_id: str,
31+
detail: str,
32+
log_message: str,
33+
log_extra: dict[str, Any] | None = None,
34+
) -> Any:
35+
"""Fetch an entity by UUID and raise HTTP 404 when absent."""
36+
result = await db.execute(select(model).filter(model.id == entity_id))
37+
entity = result.scalar_one_or_none()
38+
if entity is None:
39+
logger.warning(log_message, extra=log_extra or {"entity_id": entity_id})
40+
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=detail)
41+
return entity
42+
43+
44+
async def get_company_or_400(*, db: AsyncSession, company_id: str, actor_id: str | None = None) -> Company:
45+
"""Fetch a company by UUID and raise HTTP 400 when absent."""
46+
result = await db.execute(select(Company).filter(Company.id == company_id))
47+
company = result.scalar_one_or_none()
48+
if company is None:
49+
logger.warning(
50+
"Company not found",
51+
extra={"company_id": company_id, "actor": actor_id},
52+
)
53+
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Company not found")
54+
return company

0 commit comments

Comments
 (0)