From 3b7606107b6c430c4adc0c9e9453246f18954bb7 Mon Sep 17 00:00:00 2001 From: hexdev Date: Tue, 18 Aug 2026 19:56:14 +0700 Subject: [PATCH] feat: service token introduction --- .../551c0ad261cd_service_signatures.py | 41 +++++++++++++++++++ core/auth/fetch_sub.py | 17 ++++++++ core/auth/jwt.py | 38 +++++++++++++++-- core/secrets.py | 20 ++++++--- db/models/__init__.py | 2 + db/models/service_signatures.py | 26 ++++++++++++ repositories/service_signatures.py | 11 +++++ routes/users.py | 7 +++- schemas/dto.py | 2 +- schemas/enums.py | 5 +++ schemas/jwt.py | 13 +++++- schemas/user.py | 2 +- 12 files changed, 170 insertions(+), 14 deletions(-) create mode 100644 alembic/versions/551c0ad261cd_service_signatures.py create mode 100644 core/auth/fetch_sub.py create mode 100644 db/models/service_signatures.py create mode 100644 repositories/service_signatures.py diff --git a/alembic/versions/551c0ad261cd_service_signatures.py b/alembic/versions/551c0ad261cd_service_signatures.py new file mode 100644 index 0000000..689e37d --- /dev/null +++ b/alembic/versions/551c0ad261cd_service_signatures.py @@ -0,0 +1,41 @@ +"""+service_signatures + +Revision ID: 551c0ad261cd +Revises: cc6625f7dd7f +Create Date: 2026-08-18 19:17:33.861149 + +""" +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision: str = '551c0ad261cd' +down_revision: Union[str, Sequence[str], None] = 'cc6625f7dd7f' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + """Upgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + op.create_table('service_signatures', + sa.Column('kid', sa.TEXT(), nullable=False), + sa.Column('public_key', sa.TEXT(), nullable=False), + sa.Column('service_name', sa.TEXT(), nullable=True), + sa.Column('status', sa.Enum('ACTIVE', 'INACTIVE', name='servicesignaturestatus'), nullable=False), + sa.Column('created_at', sa.DateTime(timezone=True), nullable=False), + sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True), + sa.PrimaryKeyConstraint('kid'), + sa.UniqueConstraint('kid') + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + """Downgrade schema.""" + # ### commands auto generated by Alembic - please adjust! ### + op.drop_table('service_signatures') + # ### end Alembic commands ### diff --git a/core/auth/fetch_sub.py b/core/auth/fetch_sub.py new file mode 100644 index 0000000..687e4a9 --- /dev/null +++ b/core/auth/fetch_sub.py @@ -0,0 +1,17 @@ +from sqlalchemy.ext.asyncio import AsyncSession + +from db.models import User +from repositories.users import UserRepository +from schemas.jwt import ServiceJWTPayload + + +async def fetch_subject_from_service( + payload: ServiceJWTPayload, session: AsyncSession +) -> User | None: + repo = UserRepository(session) + + if payload.acting_as.startswith("telegram:"): + telegram_id = int(payload.acting_as.split("telegram:")[1]) + return await repo.get_user_by_telegram_id(telegram_id) + + return diff --git a/core/auth/jwt.py b/core/auth/jwt.py index 206b0f9..b12bb8e 100644 --- a/core/auth/jwt.py +++ b/core/auth/jwt.py @@ -4,17 +4,47 @@ from fastapi import HTTPException from pydantic import ValidationError from sqlalchemy.ext.asyncio import AsyncSession -from core.secrets import decode_jwt +from core.auth.fetch_sub import fetch_subject_from_service +from core.secrets import decode_jwt, decode_user_jwt, get_kid_from_token +from repositories.service_signatures import get_active_signature_by_kid from repositories.users import UserRepository from schemas.dto import AuthContext -from schemas.jwt import JWTPayload +from schemas.jwt import ServiceJWTPayload, UserJWTPayload + + +async def authorize_bot(kid: str, token: str, session: AsyncSession) -> AuthContext: + signature = await get_active_signature_by_kid(session, kid) + + if not signature: + raise HTTPException(401, detail="Invalid service signature") + + content = decode_jwt(token, signature.public_key, algo="EdDSA") + + try: + payload = ServiceJWTPayload.model_validate(content) + except ValidationError: + raise HTTPException(status_code=401, detail="Invalid credentials") from None + + if payload.exp < datetime.now(UTC).timestamp(): + raise HTTPException(status_code=401, detail="Access token expired") + + subject = await fetch_subject_from_service(payload, session) + if not subject: + raise HTTPException(status_code=401, detail="User not found") + + return AuthContext(subject, auth_method="service", service=kid) async def authorize(token: str, session: AsyncSession, service: str | None = None) -> AuthContext: - content = decode_jwt(token) + kid = get_kid_from_token(token) + + if kid: + return await authorize_bot(kid, token, session) + + content = decode_user_jwt(token) try: - payload = JWTPayload.model_validate(content) + payload = UserJWTPayload.model_validate(content) except ValidationError: raise HTTPException(status_code=401, detail="Invalid credentials") from None diff --git a/core/secrets.py b/core/secrets.py index 378ded5..14f586c 100644 --- a/core/secrets.py +++ b/core/secrets.py @@ -1,7 +1,7 @@ import hashlib import logging import secrets -from typing import Any +from typing import Any, Literal import jwt from argon2 import PasswordHasher @@ -10,7 +10,7 @@ from zxcvbn import zxcvbn from config import cfg from schemas.dto import KeyPair -from schemas.jwt import JWTPayload +from schemas.jwt import UserJWTPayload from schemas.providers import ProvidersType ctx = PasswordHasher() @@ -36,15 +36,25 @@ def generate_jwt(payload: dict[str, Any]) -> str: return jwt.encode(payload, cfg.private_key, "EdDSA") -def decode_jwt(token: str) -> dict[str, Any] | None: +def get_kid_from_token(token: str) -> str | None: + return jwt.get_unverified_header(token).get("kid") + + +def decode_jwt( + token: str, public_key: str, algo: Literal["EdDSA"] = "EdDSA" +) -> dict[str, Any] | None: try: - return jwt.decode(token, cfg.public_key, "EdDSA") + return jwt.decode(token, public_key, algo) except jwt.ExpiredSignatureError: return +def decode_user_jwt(token: str) -> dict[str, Any] | None: + return decode_jwt(token, cfg.public_key, algo="EdDSA") + + def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair: - payload = JWTPayload(sub=str(user_id), iss=iss) + payload = UserJWTPayload(sub=str(user_id), iss=iss) access_token = generate_jwt(payload.model_dump()) refresh_token = secrets.token_urlsafe(32) diff --git a/db/models/__init__.py b/db/models/__init__.py index 0b9ee52..cb3fbbb 100644 --- a/db/models/__init__.py +++ b/db/models/__init__.py @@ -2,6 +2,7 @@ from .addons import Addon from .invoice import Invoice from .orders import Order, OrderAddon from .pricing import PricingConfig +from .service_signatures import ServiceSignature from .sessions import Session from .subscription_addons import SubscriptionAddon from .subscriptions import Subscription @@ -15,6 +16,7 @@ __all__ = [ "Order", "OrderAddon", "PricingConfig", + "ServiceSignature", "Session", "Subscription", "SubscriptionAddon", diff --git a/db/models/service_signatures.py b/db/models/service_signatures.py new file mode 100644 index 0000000..2fd79e5 --- /dev/null +++ b/db/models/service_signatures.py @@ -0,0 +1,26 @@ +from datetime import UTC, datetime + +from sqlalchemy import TEXT, DateTime, Enum +from sqlalchemy.orm import Mapped, mapped_column + +from db.base import Base +from schemas.enums import ServiceSignatureStatus + + +class ServiceSignature(Base): + __tablename__ = "service_signatures" + + kid: Mapped[str] = mapped_column(TEXT, unique=True, nullable=False, primary_key=True) + public_key: Mapped[str] = mapped_column(TEXT, nullable=False) + service_name: Mapped[str] = mapped_column(TEXT, nullable=True) + status: Mapped[ServiceSignatureStatus] = mapped_column( + Enum(ServiceSignatureStatus, name="servicesignaturestatus"), + nullable=False, + default=ServiceSignatureStatus.INACTIVE, + ) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(UTC), + ) + revoked_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True) diff --git a/repositories/service_signatures.py b/repositories/service_signatures.py new file mode 100644 index 0000000..04ad38e --- /dev/null +++ b/repositories/service_signatures.py @@ -0,0 +1,11 @@ +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from db.models.service_signatures import ServiceSignature + + +async def get_active_signature_by_kid(session: AsyncSession, kid: str) -> ServiceSignature | None: + stmt = select(ServiceSignature).where(ServiceSignature.kid == kid) + r = await session.execute(stmt) + + return r.scalar_one_or_none() diff --git a/routes/users.py b/routes/users.py index 13a0100..ac81d24 100644 --- a/routes/users.py +++ b/routes/users.py @@ -12,6 +12,7 @@ from schemas.user import SubscriptionData, UserInfo router = APIRouter(prefix="/users") logger = logging.getLogger(__name__) + @router.get("/me", response_model=UserInfo) async def get_me(ctx: AuthContext = Depends(get_auth_context)): return UserInfo( @@ -37,7 +38,11 @@ async def get_subscription(ctx: AuthContext = Depends(get_auth_context)): rw_user = await get_rw_user(get_sdk(), ctx.user.telegram_id, ctx.user.username) if not rw_user: - logger.critical("RW user not found for existing local subscription (user_id=%d, sub_id=%d)", ctx.user.id, sub.id) + logger.critical( + "RW user not found for existing local subscription (user_id=%d, sub_id=%d)", + ctx.user.id, + sub.id, + ) return SubscriptionData( has_subscription=True, devices=sub.devices, diff --git a/schemas/dto.py b/schemas/dto.py index be79831..d62e095 100644 --- a/schemas/dto.py +++ b/schemas/dto.py @@ -14,5 +14,5 @@ class KeyPair: @dataclass class AuthContext: user: User - auth_method: Literal["jwt"] + auth_method: Literal["jwt", "service"] service: str | None = None diff --git a/schemas/enums.py b/schemas/enums.py index dd598f5..57751ec 100644 --- a/schemas/enums.py +++ b/schemas/enums.py @@ -4,3 +4,8 @@ from enum import StrEnum class SubscriptionStatus(StrEnum): ACTIVE = "active" EXPIRED = "expired" + + +class ServiceSignatureStatus(StrEnum): + ACTIVE = "active" + INACTIVE = "inactive" diff --git a/schemas/jwt.py b/schemas/jwt.py index 539d985..ec738b5 100644 --- a/schemas/jwt.py +++ b/schemas/jwt.py @@ -3,8 +3,17 @@ from pydantic import BaseModel, Field from core.exp import get_exp from schemas.providers import ProvidersType +type NumericDate = int | float -class JWTPayload(BaseModel): + +class UserJWTPayload(BaseModel): sub: str = Field(description="User ID") iss: ProvidersType = Field(description="Issuer") - exp: float = Field(default_factory=get_exp) + exp: NumericDate = Field(default_factory=get_exp) + + +class ServiceJWTPayload(BaseModel): + iss: str = Field(description="Service name") + iat: NumericDate = Field(description="Issued at") + exp: NumericDate = Field(description="Expiration date") + acting_as: str = Field(description="Subject of issued call, formatted as `service:id`") diff --git a/schemas/user.py b/schemas/user.py index 5d4800b..ab0fc91 100644 --- a/schemas/user.py +++ b/schemas/user.py @@ -14,6 +14,6 @@ class SubscriptionData(BaseModel): class UserInfo(BaseModel): username: str | None = Field(None) - telegram_id: str | None = Field(None) + telegram_id: int | None = Field(None) referal_code: str = Field() bonus_balance: float = Field()