feat: service token introduction

This commit is contained in:
2026-08-18 19:56:14 +07:00
parent 3d79ffb384
commit 3b7606107b
12 changed files with 170 additions and 14 deletions

View File

@@ -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 ###

17
core/auth/fetch_sub.py Normal file
View File

@@ -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

View File

@@ -4,17 +4,47 @@ from fastapi import HTTPException
from pydantic import ValidationError from pydantic import ValidationError
from sqlalchemy.ext.asyncio import AsyncSession 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 repositories.users import UserRepository
from schemas.dto import AuthContext 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: 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: try:
payload = JWTPayload.model_validate(content) payload = UserJWTPayload.model_validate(content)
except ValidationError: except ValidationError:
raise HTTPException(status_code=401, detail="Invalid credentials") from None raise HTTPException(status_code=401, detail="Invalid credentials") from None

View File

@@ -1,7 +1,7 @@
import hashlib import hashlib
import logging import logging
import secrets import secrets
from typing import Any from typing import Any, Literal
import jwt import jwt
from argon2 import PasswordHasher from argon2 import PasswordHasher
@@ -10,7 +10,7 @@ from zxcvbn import zxcvbn
from config import cfg from config import cfg
from schemas.dto import KeyPair from schemas.dto import KeyPair
from schemas.jwt import JWTPayload from schemas.jwt import UserJWTPayload
from schemas.providers import ProvidersType from schemas.providers import ProvidersType
ctx = PasswordHasher() ctx = PasswordHasher()
@@ -36,15 +36,25 @@ def generate_jwt(payload: dict[str, Any]) -> str:
return jwt.encode(payload, cfg.private_key, "EdDSA") 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: try:
return jwt.decode(token, cfg.public_key, "EdDSA") return jwt.decode(token, public_key, algo)
except jwt.ExpiredSignatureError: except jwt.ExpiredSignatureError:
return 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: 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()) access_token = generate_jwt(payload.model_dump())
refresh_token = secrets.token_urlsafe(32) refresh_token = secrets.token_urlsafe(32)

View File

@@ -2,6 +2,7 @@ from .addons import Addon
from .invoice import Invoice from .invoice import Invoice
from .orders import Order, OrderAddon from .orders import Order, OrderAddon
from .pricing import PricingConfig from .pricing import PricingConfig
from .service_signatures import ServiceSignature
from .sessions import Session from .sessions import Session
from .subscription_addons import SubscriptionAddon from .subscription_addons import SubscriptionAddon
from .subscriptions import Subscription from .subscriptions import Subscription
@@ -15,6 +16,7 @@ __all__ = [
"Order", "Order",
"OrderAddon", "OrderAddon",
"PricingConfig", "PricingConfig",
"ServiceSignature",
"Session", "Session",
"Subscription", "Subscription",
"SubscriptionAddon", "SubscriptionAddon",

View File

@@ -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)

View File

@@ -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()

View File

@@ -12,6 +12,7 @@ from schemas.user import SubscriptionData, UserInfo
router = APIRouter(prefix="/users") router = APIRouter(prefix="/users")
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@router.get("/me", response_model=UserInfo) @router.get("/me", response_model=UserInfo)
async def get_me(ctx: AuthContext = Depends(get_auth_context)): async def get_me(ctx: AuthContext = Depends(get_auth_context)):
return UserInfo( 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) rw_user = await get_rw_user(get_sdk(), ctx.user.telegram_id, ctx.user.username)
if not rw_user: 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( return SubscriptionData(
has_subscription=True, has_subscription=True,
devices=sub.devices, devices=sub.devices,

View File

@@ -14,5 +14,5 @@ class KeyPair:
@dataclass @dataclass
class AuthContext: class AuthContext:
user: User user: User
auth_method: Literal["jwt"] auth_method: Literal["jwt", "service"]
service: str | None = None service: str | None = None

View File

@@ -4,3 +4,8 @@ from enum import StrEnum
class SubscriptionStatus(StrEnum): class SubscriptionStatus(StrEnum):
ACTIVE = "active" ACTIVE = "active"
EXPIRED = "expired" EXPIRED = "expired"
class ServiceSignatureStatus(StrEnum):
ACTIVE = "active"
INACTIVE = "inactive"

View File

@@ -3,8 +3,17 @@ from pydantic import BaseModel, Field
from core.exp import get_exp from core.exp import get_exp
from schemas.providers import ProvidersType from schemas.providers import ProvidersType
type NumericDate = int | float
class JWTPayload(BaseModel):
class UserJWTPayload(BaseModel):
sub: str = Field(description="User ID") sub: str = Field(description="User ID")
iss: ProvidersType = Field(description="Issuer") 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`")

View File

@@ -14,6 +14,6 @@ class SubscriptionData(BaseModel):
class UserInfo(BaseModel): class UserInfo(BaseModel):
username: str | None = Field(None) username: str | None = Field(None)
telegram_id: str | None = Field(None) telegram_id: int | None = Field(None)
referal_code: str = Field() referal_code: str = Field()
bonus_balance: float = Field() bonus_balance: float = Field()