feat: service token introduction
This commit is contained in:
41
alembic/versions/551c0ad261cd_service_signatures.py
Normal file
41
alembic/versions/551c0ad261cd_service_signatures.py
Normal 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
17
core/auth/fetch_sub.py
Normal 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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
26
db/models/service_signatures.py
Normal file
26
db/models/service_signatures.py
Normal 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)
|
||||||
11
repositories/service_signatures.py
Normal file
11
repositories/service_signatures.py
Normal 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()
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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`")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user