Skip to content

Commit 4b57907

Browse files
committed
added method to get claims
1 parent 7584d34 commit 4b57907

4 files changed

Lines changed: 67 additions & 10 deletions

File tree

api-server/app/auth/controllers/create_token.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55

66
from ..models.token_request import TokenRequest
77
from ..models.token_response import TokenResponse
8+
from ..models.token_claims import TokenClaims
89

910
from app.singletons.logs_manager import LogsManager
1011

@@ -38,17 +39,17 @@ async def create_token(request: TokenRequest, x_exosphere_request_id: str) -> To
3839

3940
logger.info("User is a super admin", x_exosphere_request_id=x_exosphere_request_id)
4041

41-
token_claims = {
42-
"user_id": str(user.id),
43-
"user_name": user.name,
44-
"user_type": user.type,
45-
"verification_status": user.verification_status,
46-
"status": user.status,
47-
"exp": datetime.now() + timedelta(seconds=JWT_EXPIRES_IN)
48-
}
42+
token_claims = TokenClaims(
43+
user_id=str(user.id),
44+
user_name=user.name,
45+
user_type=user.type,
46+
verification_status=user.verification_status,
47+
status=user.status,
48+
exp=int((datetime.now() + timedelta(seconds=JWT_EXPIRES_IN)).timestamp())
49+
)
4950

5051
return TokenResponse(
51-
access_token=jwt.encode(token_claims, JWT_SECRET_KEY, algorithm=JWT_ALGORITHM)
52+
access_token=jwt.encode(token_claims.model_dump(), JWT_SECRET_KEY, algorithm=JWT_ALGORITHM)
5253
)
5354

5455
except Exception as e:
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
import jwt
2+
import os
3+
4+
from starlette.middleware.base import BaseHTTPMiddleware
5+
from app.singletons.logs_manager import LogsManager
6+
from starlette.requests import Request
7+
from starlette.responses import JSONResponse
8+
9+
from ..models.token_claims import TokenClaims
10+
11+
logger = LogsManager().get_logger()
12+
13+
JWT_SECRET_KEY = os.getenv("JWT_SECRET_KEY")
14+
if not JWT_SECRET_KEY:
15+
raise ValueError("JWT_SECRET_KEY environment variable is not set or is empty.")
16+
JWT_ALGORITHM = "HS256"
17+
18+
19+
class GetTokenClaimsMiddleware(BaseHTTPMiddleware):
20+
async def dispatch(self, request: Request, call_next):
21+
22+
token = request.headers.get("Authorization")
23+
24+
if token:
25+
if not token.startswith("Bearer "):
26+
logger.error("Invalid token format", x_exosphere_request_id=getattr(request.state, 'x_exosphere_request_id', None))
27+
return JSONResponse(status_code=401, content={"message": "Invalid token format", "success": False})
28+
29+
try:
30+
token_claims = jwt.decode(token.split(" ")[1], JWT_SECRET_KEY, algorithms=[JWT_ALGORITHM])
31+
request.state.token_claims = TokenClaims(**token_claims)
32+
logger.info("Token claims decoded", x_exosphere_request_id=getattr(request.state, 'x_exosphere_request_id', None), user_id=request.state.token_claims.user_id)
33+
34+
except Exception as e:
35+
logger.error("Error decoding token", error=e, x_exosphere_request_id=getattr(request.state, 'x_exosphere_request_id', None))
36+
37+
return JSONResponse(status_code=401, content={"message": "Invalid token", "success": False})
38+
39+
else:
40+
logger.error("No token provided", x_exosphere_request_id=getattr(request.state, 'x_exosphere_request_id', None))
41+
request.state.token_claims = None
42+
43+
return await call_next(request)
Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,9 @@
1+
from pydantic import BaseModel
2+
3+
class TokenClaims(BaseModel):
4+
user_id: str
5+
user_name: str
6+
user_type: str
7+
verification_status: str
8+
status: str
9+
exp: int

api-server/app/main.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
UnhandledExceptionsMiddleware,
1717
)
1818
from .middlewares.request_id_middleware import RequestIdMiddleware
19+
from .auth.middlewares.get_token_claims import GetTokenClaimsMiddleware
1920

2021
# injecting databases
2122
from .user.models.user_database_model import User
@@ -49,13 +50,16 @@ async def lifespan(app: FastAPI):
4950

5051
app = FastAPI(lifespan=lifespan)
5152

53+
5254
# this middleware should be the first one
55+
app.add_middleware(GetTokenClaimsMiddleware)
56+
5357
app.add_middleware(RequestIdMiddleware)
5458

55-
# this middleware should be the last one
5659
app.add_middleware(UnhandledExceptionsMiddleware)
5760

5861

62+
5963
@app.get("/health-check")
6064
def health() -> dict:
6165
return {"message": "OK"}

0 commit comments

Comments
 (0)