Compare commits

...

36 Commits

Author SHA1 Message Date
4262d4d017 fix: replaced remnawave_sdk to local gitea 2026-09-11 21:55:28 +07:00
429c961ace feat: /signup returns access_token 2026-09-04 21:21:20 +07:00
27e2f58956 feat: local testing compose file 2026-09-04 21:02:08 +07:00
a374501e4d fix!: removed .commit() from repository level 2026-08-24 11:07:26 +07:00
8d1b753b99 feat: advanced RW sync system 2026-08-24 10:22:14 +07:00
b39cae8046 feat: cron worker for renewal notifs 2026-08-20 21:28:43 +07:00
0a58a41930 feat: /internal/renewal/pending 2026-08-20 21:01:22 +07:00
05ded7d7ea chore: minor tweaks for better architecture 2026-08-20 11:43:48 +07:00
e212492b30 fix: rw fetching + addon orm refreshing 2026-08-20 11:22:39 +07:00
b759a997d9 fix: addon.is_enabled NOT NULL 2026-08-19 21:58:57 +07:00
77ef4aaa57 fix(auth): handle all PyJWT errors gracefully and enforce link code status checks 2026-08-19 21:12:45 +07:00
2e01b6502c feat(link_code): added tests 2026-08-18 21:21:20 +07:00
771d44d34c feat(link_codes): linking method between telegram and website 2026-08-18 21:10:57 +07:00
7ae3c98585 feat: introduced pytest powered tests 2026-08-18 20:09:59 +07:00
938e924107 chore: removed deprecated tests/ from .gitignore 2026-08-18 19:57:16 +07:00
3b7606107b feat: service token introduction 2026-08-18 19:56:14 +07:00
3d79ffb384 fix(/users): missing rw entry doesn't cause errors 2026-08-18 19:05:42 +07:00
039babf540 feat: server-side password security checks 2026-08-18 12:36:22 +07:00
9e209dd695 feat: addon.is_enabled 2026-08-17 21:57:38 +07:00
05914f0b23 feat: added duration_days to /users/subscriptions 2026-08-17 21:33:08 +07:00
8b6c43f89a fix: adjusted BillType due to 422 2026-08-17 13:03:10 +07:00
865c5cb9b3 feat: healthcheck endpoint file 2026-08-14 19:48:43 +07:00
9fa8e8c9a9 feat: healthcheck endpoint 2026-08-14 19:47:41 +07:00
9a62722a3c fix: match the BillType to actual API responses 2026-08-14 13:53:16 +07:00
5d610c484d feat: advanced catches at payment processing to prevent the payment loss 2026-08-12 12:54:30 +07:00
8cb1ac89a6 fix: fixes to match balance fields and referal balance fixes 2026-08-12 11:45:15 +07:00
d1f7612f37 feat: added balance to userinfo 2026-08-12 11:28:31 +07:00
03e6363c1a fix: strict has_subscription statements. 2026-08-10 12:17:44 +07:00
c06be7b5f2 feat(/me): subscription data + rw controls 2026-08-03 11:15:11 +07:00
801014c392 feat: referal spec on registration 2026-08-03 10:35:26 +07:00
680c15414c feat: GET /me endpoint 2026-08-03 10:23:44 +07:00
dee2cad540 refactor(payments): extract subscription purchase processing to a service 2026-08-03 10:05:21 +07:00
74e952698b chore: deleted testing files 2026-08-03 00:01:50 +07:00
1e4cb43ac7 feat(rw): remnawave integration 2026-08-03 00:01:00 +07:00
97a4e819d6 feat(sub): handling upgrade/downgrade
rw integration and /me endpoints are left
2026-08-02 14:16:57 +07:00
24f857ef1b feat(sub): subscription logic (pre release) 2026-08-02 13:23:50 +07:00
83 changed files with 3487 additions and 249 deletions

3
.gitignore vendored
View File

@@ -179,5 +179,4 @@ cython_debug/
plans
dev.sh
test_dummy.py
*.pem
tests/
*.pem

1
.python-version Normal file
View File

@@ -0,0 +1 @@
3.13.0

View File

@@ -0,0 +1,46 @@
"""+order.durationdays
Revision ID: 0ea5625b4913
Revises: b97ce5b7d663
Create Date: 2026-08-02 13:17:18.429961
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = '0ea5625b4913'
down_revision: Union[str, Sequence[str], None] = 'b97ce5b7d663'
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! ###
orderstatus = postgresql.ENUM('PENDING', 'PAID', 'DISCARDED', name='orderstatus')
orderstatus.create(op.get_bind())
op.create_unique_constraint(None, 'order_addons', ['id'])
op.add_column('orders', sa.Column('duration_days', sa.INTEGER(), nullable=False))
op.add_column('orders', sa.Column('status', sa.Enum('PENDING', 'PAID', 'DISCARDED', name='orderstatus'), nullable=False))
op.create_unique_constraint(None, 'orders', ['id'])
op.create_unique_constraint(None, 'subscription_addons', ['id'])
op.create_unique_constraint(None, 'subscriptions', ['id'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'subscriptions', type_='unique')
op.drop_constraint(None, 'subscription_addons', type_='unique')
op.drop_constraint(None, 'orders', type_='unique')
op.drop_column('orders', 'status')
op.drop_column('orders', 'duration_days')
op.drop_constraint(None, 'order_addons', type_='unique')
# ### end Alembic commands ###

View File

@@ -0,0 +1,34 @@
"""+user.referal_code
Revision ID: 135672cdf14a
Revises: 72e78a8a43cf
Create Date: 2026-08-03 10:26:22.290547
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '135672cdf14a'
down_revision: Union[str, Sequence[str], None] = '72e78a8a43cf'
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.add_column('users', sa.Column('referal_code', sa.TEXT(), nullable=False))
op.create_unique_constraint(None, 'users', ['referal_code'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'users', type_='unique')
op.drop_column('users', 'referal_code')
# ### end Alembic commands ###

View File

@@ -0,0 +1,46 @@
"""add rw sync outbox
Revision ID: 426b372286c3
Revises: af1c3d7e4b20
Create Date: 2026-08-20 21:31:20.949847
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '426b372286c3'
down_revision: Union[str, Sequence[str], None] = 'af1c3d7e4b20'
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('rw_sync_outbox',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('subscription_id', sa.INTEGER(), nullable=False),
sa.Column('revision', sa.INTEGER(), nullable=False),
sa.Column('attempts', sa.INTEGER(), nullable=False),
sa.Column('next_attempt_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('locked_until', sa.DateTime(timezone=True), nullable=True),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['subscription_id'], ['subscriptions.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('subscription_id')
)
op.create_unique_constraint(None, 'service_notifications', ['id'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'service_notifications', type_='unique')
op.drop_table('rw_sync_outbox')
# ### end Alembic commands ###

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

View File

@@ -0,0 +1,44 @@
"""+link_codes
Revision ID: 6d1bae3dc723
Revises: 551c0ad261cd
Create Date: 2026-08-18 20:19:35.961227
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '6d1bae3dc723'
down_revision: Union[str, Sequence[str], None] = '551c0ad261cd'
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('link_codes',
sa.Column('code', sa.TEXT(), nullable=False),
sa.Column('user_id', sa.INTEGER(), nullable=False),
sa.Column('status', sa.Enum('ACTIVE', 'USED', 'EXPIRED', name='linkcodestatus'), nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('used_at', sa.DateTime(timezone=True), nullable=True),
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ),
sa.PrimaryKeyConstraint('code'),
sa.UniqueConstraint('code')
)
op.create_unique_constraint(None, 'service_signatures', ['kid'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'service_signatures', type_='unique')
op.drop_table('link_codes')
# ### end Alembic commands ###

View File

@@ -0,0 +1,42 @@
"""order financials and invoice link
Revision ID: 72e78a8a43cf
Revises: bc424721d767
Create Date: 2026-08-02 13:55:18.013052
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '72e78a8a43cf'
down_revision: Union[str, Sequence[str], None] = 'bc424721d767'
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.add_column('invoices', sa.Column('order_id', sa.INTEGER(), nullable=True))
op.create_foreign_key(None, 'invoices', 'orders', ['order_id'], ['id'])
op.add_column('orders', sa.Column('total_amount', sa.FLOAT(), nullable=False))
op.add_column('orders', sa.Column('balance_amount', sa.FLOAT(), nullable=False))
op.add_column('orders', sa.Column('applies_at', sa.DateTime(timezone=True), nullable=True))
op.add_column('orders', sa.Column('applied_at', sa.DateTime(timezone=True), nullable=True))
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('orders', 'applied_at')
op.drop_column('orders', 'applies_at')
op.drop_column('orders', 'balance_amount')
op.drop_column('orders', 'total_amount')
op.drop_constraint(None, 'invoices', type_='foreignkey')
op.drop_column('invoices', 'order_id')
# ### end Alembic commands ###

View File

@@ -0,0 +1,38 @@
"""enforce addon.is_enabled->NOT NULL
Revision ID: ad537d63c440
Revises: 6d1bae3dc723
Create Date: 2026-08-19 21:14:39.928563
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'ad537d63c440'
down_revision: Union[str, Sequence[str], None] = '6d1bae3dc723'
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.alter_column('addons', 'is_enabled',
existing_type=sa.BOOLEAN(),
nullable=False)
op.create_unique_constraint(None, 'link_codes', ['code'])
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_constraint(None, 'link_codes', type_='unique')
op.alter_column('addons', 'is_enabled',
existing_type=sa.BOOLEAN(),
nullable=True)
# ### end Alembic commands ###

View File

@@ -0,0 +1,32 @@
"""Remove global service notification expiry uniqueness.
Revision ID: af1c3d7e4b20
Revises: ee6174eeef14
Create Date: 2026-08-20 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
revision: str = "af1c3d7e4b20"
down_revision: Union[str, Sequence[str], None] = "ee6174eeef14"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.drop_constraint(
"service_notifications_sub_expires_at_key",
"service_notifications",
type_="unique",
)
def downgrade() -> None:
op.create_unique_constraint(
"service_notifications_sub_expires_at_key",
"service_notifications",
["sub_expires_at"],
)

View File

@@ -0,0 +1,72 @@
"""subscriptions infra
Revision ID: b97ce5b7d663
Revises: 3d767875ec7d
Create Date: 2026-08-02 12:38:25.357158
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'b97ce5b7d663'
down_revision: Union[str, Sequence[str], None] = '3d767875ec7d'
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('orders',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('user_id', sa.INTEGER(), nullable=False),
sa.Column('devices', sa.INTEGER(), nullable=False),
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id'),
sa.UniqueConstraint('user_id')
)
op.create_table('subscriptions',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('user_id', sa.INTEGER(), nullable=False),
sa.Column('devices', sa.INTEGER(), nullable=False),
sa.Column('status', sa.Enum('ACTIVE', 'EXPIRED', name='subscriptionstatus'), nullable=False),
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id'),
sa.UniqueConstraint('user_id')
)
op.create_table('order_addons',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('order_id', sa.INTEGER(), nullable=False),
sa.Column('addon_id', sa.TEXT(), nullable=False),
sa.ForeignKeyConstraint(['addon_id'], ['addons.id'], ),
sa.ForeignKeyConstraint(['order_id'], ['orders.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id')
)
op.create_table('subscription_addons',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('subscription_id', sa.INTEGER(), nullable=False),
sa.Column('addon_id', sa.TEXT(), nullable=False),
sa.ForeignKeyConstraint(['addon_id'], ['addons.id'], ),
sa.ForeignKeyConstraint(['subscription_id'], ['subscriptions.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id')
)
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_table('subscription_addons')
op.drop_table('order_addons')
op.drop_table('subscriptions')
op.drop_table('orders')
# ### end Alembic commands ###

View File

@@ -0,0 +1,32 @@
"""order.user_id -> NOT UNIQUE
Revision ID: bc424721d767
Revises: 0ea5625b4913
Create Date: 2026-08-02 13:23:23.626697
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'bc424721d767'
down_revision: Union[str, Sequence[str], None] = '0ea5625b4913'
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.drop_constraint(op.f('orders_user_id_key'), 'orders', type_='unique')
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.create_unique_constraint(op.f('orders_user_id_key'), 'orders', ['user_id'], postgresql_nulls_not_distinct=False)
# ### end Alembic commands ###

View File

@@ -22,7 +22,7 @@ def upgrade() -> None:
"""Upgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.create_unique_constraint(None, 'invoices', ['id'])
op.add_column('users', sa.Column('balance', sa.FLOAT(precision=2), nullable=False))
op.add_column('users', sa.Column('balance', sa.FLOAT(precision=2), nullable=False, default=0))
# ### end Alembic commands ###

View File

@@ -0,0 +1,32 @@
"""+addon.is_enabled
Revision ID: cc6625f7dd7f
Revises: f4ece5936e3d
Create Date: 2026-08-17 21:39:00.054081
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'cc6625f7dd7f'
down_revision: Union[str, Sequence[str], None] = 'f4ece5936e3d'
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.add_column('addons', sa.Column('is_enabled', sa.BOOLEAN(), nullable=True))
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('addons', 'is_enabled')
# ### end Alembic commands ###

View File

@@ -0,0 +1,47 @@
"""+service_notifications
Revision ID: ee6174eeef14
Revises: ad537d63c440
Create Date: 2026-08-20 12:01:51.073160
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'ee6174eeef14'
down_revision: Union[str, Sequence[str], None] = 'ad537d63c440'
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_notifications',
sa.Column('id', sa.INTEGER(), autoincrement=True, nullable=False),
sa.Column('subscription_id', sa.INTEGER(), nullable=False),
sa.Column('notify_type', sa.Enum('SEVEN_DAYS', 'THREE_DAYS', 'ONE_DAY', 'EXPIRED', name='notificationtype'), nullable=False),
sa.Column('sub_expires_at', sa.DateTime(timezone=True), nullable=False),
sa.Column('status', sa.Enum('PENDING', 'DISPATCHED', 'SENT', 'FAILED', name='notificationstatus'), nullable=False),
sa.Column('dispatched_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('sent_at', sa.DateTime(timezone=True), nullable=True),
sa.Column('attempts', sa.INTEGER(), nullable=False),
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(['subscription_id'], ['subscriptions.id'], ),
sa.PrimaryKeyConstraint('id'),
sa.UniqueConstraint('id'),
sa.UniqueConstraint('sub_expires_at'),
sa.UniqueConstraint('subscription_id', 'notify_type', 'sub_expires_at', name='uq_subscription_notification')
)
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_table('service_notifications')
# ### end Alembic commands ###

View File

@@ -0,0 +1,32 @@
"""+subscription.duration_days
Revision ID: f4ece5936e3d
Revises: 135672cdf14a
Create Date: 2026-08-17 21:23:17.900314
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'f4ece5936e3d'
down_revision: Union[str, Sequence[str], None] = '135672cdf14a'
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.add_column('subscriptions', sa.Column('duration_days', sa.INTEGER(), nullable=True))
# ### end Alembic commands ###
def downgrade() -> None:
"""Downgrade schema."""
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('subscriptions', 'duration_days')
# ### end Alembic commands ###

4
compose.local.yml Normal file
View File

@@ -0,0 +1,4 @@
services:
postgres:
ports: !override
- "5432:5432"

View File

@@ -5,21 +5,50 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env")
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
### Internal settings ###
postgres_user: str = Field()
postgres_password: str = Field()
postgres_host: str = Field()
postgres_port: str = Field()
postgres_db: str = Field()
### Telegram ###
bot_username: str = Field()
### Remnawave ###
remnawave_base_url: str = Field()
remnawave_sub_url: str = Field()
remnawave_token: str = Field()
remnawave_default_squads_raw: str = Field(alias="REMNAWAVE_DEFAULT_SQUADS_UUIDS")
### Finance ###
minimal_deposit: int = Field()
referal_bonus: int = Field(30)
#### Pally ####
pally_shop_id: str = Field()
pally_token: str = Field()
### Security related ###
private_key_fp: str = Field()
public_key_fp: str = Field()
access_token_ttl: int = Field(description="Access token TTL (minutes)")
link_code_ttl: int = Field(8, description="Link Code TTL (minutes)")
notification_scan_interval: int = Field(5, ge=1)
rw_sync_interval_minutes: int = Field(1, ge=1)
rw_sync_reconcile_interval_minutes: int = Field(60, ge=1)
pally_shop_id: str = Field()
pally_token: str = Field()
min_password_length: int = Field(8)
link_code_length: int = Field(8)
password_security_threshold: int = Field(2)
@computed_field
@property
def remnawave_default_squads(self) -> list[str]:
return self.remnawave_default_squads_raw.split(",")
@computed_field
@property
@@ -42,5 +71,10 @@ class Settings(BaseSettings):
with open(self.public_key_fp, "rb") as f:
return f.read()
@computed_field
@property
def bot_url(self) -> str:
return "https://t.me/" + self.bot_username.lstrip("@")
cfg = Settings() # type: ignore

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

@@ -0,0 +1,14 @@
from db.models import User
from db.session import UnitOfWork
from repositories.users import UserRepository
from schemas.jwt import ServiceJWTPayload
async def fetch_subject_from_service(payload: ServiceJWTPayload, uow: UnitOfWork) -> User | None:
repo = UserRepository(uow)
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

@@ -2,26 +2,56 @@ from datetime import UTC, datetime
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 db.session import UnitOfWork
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(token: str, session: AsyncSession, service: str | None = None) -> AuthContext:
content = decode_jwt(token)
async def authorize_bot(kid: str, token: str, uow: UnitOfWork) -> AuthContext:
signature = await get_active_signature_by_kid(uow.session, kid)
if not signature:
raise HTTPException(401, detail="Invalid service signature")
content = decode_jwt(token, signature.public_key, algo="EdDSA")
try:
payload = JWTPayload.model_validate(content)
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")
repo = UserRepository(session)
subject = await fetch_subject_from_service(payload, uow)
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, uow: UnitOfWork, service: str | None = None) -> AuthContext:
kid = get_kid_from_token(token)
if kid:
return await authorize_bot(kid, token, uow)
content = decode_user_jwt(token)
try:
payload = UserJWTPayload.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")
repo = UserRepository(uow)
user = await repo.get_user_by_id(int(payload.sub))
if not user:

View File

@@ -1,15 +1,17 @@
from fastapi import Depends, HTTPException, Request
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.auth import jwt
from db.session import get_db
from core.secrets import get_kid_from_token
from db.session import UnitOfWork, get_uow
from external.pally import PallyClient
from schemas.dto import AuthContext
from repositories.service_signatures import get_active_signature_by_kid
from schemas.dto import AuthContext, ServiceIdentity
from services.subscriptions import sync_user_subscription
async def get_auth_context(
request: Request, session: AsyncSession = Depends(get_db)
request: Request, uow: UnitOfWork = Depends(get_uow)
) -> AuthContext | None:
auth = request.headers.get("Authorization")
@@ -18,7 +20,28 @@ async def get_auth_context(
if auth.startswith("Bearer"):
token = auth.removeprefix("Bearer ").strip()
return await jwt.authorize(token, session)
ctx = await jwt.authorize(token, uow)
await sync_user_subscription(uow.session, user=ctx.user)
await uow.commit()
return ctx
async def get_service_identity(
request: Request, uow: UnitOfWork = Depends(get_uow)
) -> ServiceIdentity | None:
auth = request.headers.get("Authorization")
if not auth:
raise HTTPException(401)
if auth.startswith("Bearer"):
token = auth.removeprefix("Bearer ").strip()
kid = get_kid_from_token(token)
if not kid:
raise HTTPException(401, detail="No kid provided.")
signature = await get_active_signature_by_kid(uow.session, kid)
if not signature:
raise HTTPException(401, detail="Invalid signature")
return ServiceIdentity(service=signature.kid)
def get_pally_client() -> PallyClient:

View File

@@ -1,15 +1,16 @@
import hashlib
import logging
import secrets
from typing import Any
from typing import Any, Literal
import jwt
from argon2 import PasswordHasher
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
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()
@@ -35,15 +36,29 @@ 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:
try:
return jwt.decode(token, cfg.public_key, "EdDSA")
except jwt.ExpiredSignatureError:
header = jwt.get_unverified_header(token)
return header.get("kid")
except jwt.exceptions.PyJWTError:
return
def decode_jwt(
token: str, public_key: str, algo: Literal["EdDSA"] = "EdDSA"
) -> dict[str, Any] | None:
try:
return jwt.decode(token, public_key, algo)
except jwt.exceptions.PyJWTError:
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)
@@ -52,3 +67,11 @@ def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair:
def hash_refresh_token(token: str):
return hashlib.sha256(token.encode()).hexdigest()
def estimate_password_strength(password: str) -> bool:
if len(password) < cfg.min_password_length:
return False
r = zxcvbn(password)
return r.get("score", 0) > cfg.password_security_threshold

View File

@@ -1,8 +1,30 @@
from .addons import Addon
from .invoice import Invoice
from .link_codes import LinkCode
from .orders import Order, OrderAddon
from .pricing import PricingConfig
from .rw_sync_outbox import RWSyncOutbox
from .service_notifications import ServiceNotification
from .service_signatures import ServiceSignature
from .sessions import Session
from .subscription_addons import SubscriptionAddon
from .subscriptions import Subscription
from .transactions import BalanceTransaction
from .users import User
__all__ = ["Addon", "BalanceTransaction", "Invoice", "PricingConfig", "Session", "User"]
__all__ = [
"Addon",
"BalanceTransaction",
"Invoice",
"LinkCode",
"Order",
"OrderAddon",
"PricingConfig",
"RWSyncOutbox",
"ServiceNotification",
"ServiceSignature",
"Session",
"Subscription",
"SubscriptionAddon",
"User",
]

View File

@@ -1,4 +1,4 @@
from sqlalchemy import FLOAT, INTEGER, TEXT
from sqlalchemy import BOOLEAN, FLOAT, INTEGER, TEXT
from sqlalchemy.orm import Mapped, mapped_column
from db.base import Base
@@ -11,3 +11,4 @@ class Addon(Base):
name: Mapped[str] = mapped_column(TEXT, nullable=False, unique=False)
price: Mapped[float] = mapped_column(FLOAT, nullable=False)
free_threshold: Mapped[int] = mapped_column(INTEGER, default=-1)
is_enabled: Mapped[bool] = mapped_column(BOOLEAN, default=False, nullable=False)

View File

@@ -7,7 +7,7 @@ from db.base import Base
from schemas.invoices import InvoiceStatus
if TYPE_CHECKING:
from db.models import User
from db.models import Order, User
class Invoice(Base):
@@ -17,9 +17,11 @@ class Invoice(Base):
INTEGER, autoincrement=True, unique=True, nullable=False, primary_key=True
)
creator_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False)
order_id: Mapped[int | None] = mapped_column(ForeignKey("orders.id"), nullable=True)
amount: Mapped[float] = mapped_column(FLOAT, nullable=False)
status: Mapped[InvoiceStatus] = mapped_column(
Enum(InvoiceStatus, name="invoicestatus"), nullable=False, default=InvoiceStatus.ACTIVE
)
creator: Mapped["User"] = relationship("User", lazy="selectin")
order: Mapped["Order | None"] = relationship("Order", lazy="selectin")

28
db/models/link_codes.py Normal file
View File

@@ -0,0 +1,28 @@
from datetime import UTC, datetime
from sqlalchemy import TEXT, DateTime, Enum, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column
from db.base import Base
from schemas.enums import LinkCodeStatus
class LinkCode(Base):
__tablename__ = "link_codes"
code: Mapped[str] = mapped_column(TEXT, nullable=False, unique=True, primary_key=True)
user_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False)
status: Mapped[LinkCodeStatus] = mapped_column(
Enum(LinkCodeStatus, name="linkcodestatus"), nullable=False, default=LinkCodeStatus.ACTIVE
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,
default=lambda: datetime.now(UTC),
)
expires_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,
)
used_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)

53
db/models/orders.py Normal file
View File

@@ -0,0 +1,53 @@
import datetime
from enum import StrEnum
from typing import TYPE_CHECKING
from sqlalchemy import FLOAT, INTEGER, DateTime, Enum, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
if TYPE_CHECKING:
from db.models import User
class OrderStatus(StrEnum):
PENDING = "pending"
PAID = "paid"
DISCARDED = "discarded"
class OrderAddon(Base):
__tablename__ = "order_addons"
id: Mapped[int] = mapped_column(
INTEGER, nullable=False, unique=True, autoincrement=True, primary_key=True
)
order_id: Mapped[int] = mapped_column(ForeignKey("orders.id"), nullable=False)
addon_id: Mapped[str] = mapped_column(ForeignKey("addons.id"), nullable=False)
order: Mapped["Order"] = relationship("Order", back_populates="addons", lazy="selectin")
class Order(Base):
__tablename__ = "orders"
id: Mapped[int] = mapped_column(
INTEGER, nullable=False, autoincrement=True, unique=True, primary_key=True
)
user_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False)
devices: Mapped[int] = mapped_column(INTEGER, nullable=False)
duration_days: Mapped[int] = mapped_column(INTEGER, nullable=False)
total_amount: Mapped[float] = mapped_column(FLOAT, nullable=False)
balance_amount: Mapped[float] = mapped_column(FLOAT, nullable=False, default=0)
status: Mapped[OrderStatus] = mapped_column(
Enum(OrderStatus, name="orderstatus"), nullable=False, default=OrderStatus.PENDING
)
applies_at: Mapped[datetime.datetime | None] = mapped_column(DateTime(True), nullable=True)
applied_at: Mapped[datetime.datetime | None] = mapped_column(DateTime(True), nullable=True)
user: Mapped["User"] = relationship("User", back_populates="orders", lazy="selectin")
addons: Mapped[list["OrderAddon"]] = relationship(
"OrderAddon", back_populates="order", lazy="selectin"
)

View File

@@ -0,0 +1,34 @@
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from sqlalchemy import INTEGER, DateTime, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
if TYPE_CHECKING:
from db.models import Subscription
class RWSyncOutbox(Base):
__tablename__ = "rw_sync_outbox"
id: Mapped[int] = mapped_column(INTEGER, primary_key=True, autoincrement=True)
subscription_id: Mapped[int] = mapped_column(
ForeignKey("subscriptions.id"), nullable=False, unique=True
)
revision: Mapped[int] = mapped_column(INTEGER, nullable=False, default=1)
attempts: Mapped[int] = mapped_column(INTEGER, nullable=False, default=0)
next_attempt_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
locked_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC)
)
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,
default=lambda: datetime.now(UTC),
onupdate=lambda: datetime.now(UTC),
)
subscription: Mapped["Subscription"] = relationship("Subscription", lazy="selectin")

View File

@@ -0,0 +1,52 @@
from datetime import UTC, datetime
from typing import TYPE_CHECKING
from sqlalchemy import INTEGER, DateTime, Enum, ForeignKey, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
from schemas.enums import NotificationStatus, NotificationType
if TYPE_CHECKING:
from db.models import Subscription
class ServiceNotification(Base):
__tablename__ = "service_notifications"
id: Mapped[int] = mapped_column(
INTEGER, unique=True, autoincrement=True, primary_key=True, nullable=False
)
subscription_id: Mapped[int] = mapped_column(
ForeignKey("subscriptions.id"), nullable=False, unique=False
)
notify_type: Mapped[NotificationType] = mapped_column(
Enum(NotificationType, name="notificationtype"), nullable=False
)
sub_expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
status: Mapped[NotificationStatus] = mapped_column(
Enum(NotificationStatus, name="notificationstatus"),
nullable=False,
default=NotificationStatus.PENDING,
)
dispatched_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)
sent_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)
attempts: Mapped[int] = mapped_column(INTEGER, nullable=False, default=0)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,
default=lambda: datetime.now(UTC),
)
__table_args__ = (
UniqueConstraint(
"subscription_id",
"notify_type",
"sub_expires_at",
name="uq_subscription_notification",
),
)
subscription: Mapped["Subscription"] = relationship("Subscription", lazy="selectin")

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,23 @@
from typing import TYPE_CHECKING
from sqlalchemy import INTEGER, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
if TYPE_CHECKING:
from db.models import Subscription
class SubscriptionAddon(Base):
__tablename__ = "subscription_addons"
id: Mapped[int] = mapped_column(
INTEGER, nullable=False, unique=True, autoincrement=True, primary_key=True
)
subscription_id: Mapped[int] = mapped_column(ForeignKey("subscriptions.id"), nullable=False)
addon_id: Mapped[str] = mapped_column(ForeignKey("addons.id"), nullable=False)
subscription: Mapped["Subscription"] = relationship(
"Subscription", back_populates="addons", lazy="selectin"
)

View File

@@ -0,0 +1,34 @@
import datetime
from typing import TYPE_CHECKING
from sqlalchemy import INTEGER, DateTime, Enum, ForeignKey
from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
from schemas.enums import SubscriptionStatus
if TYPE_CHECKING:
from db.models import SubscriptionAddon, User
class Subscription(Base):
__tablename__ = "subscriptions"
id: Mapped[int] = mapped_column(
INTEGER, nullable=False, autoincrement=True, unique=True, primary_key=True
)
user_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False, unique=True)
devices: Mapped[int] = mapped_column(INTEGER, nullable=False)
duration_days: Mapped[int] = mapped_column(INTEGER, default=30, nullable=True)
status: Mapped[SubscriptionStatus] = mapped_column(
Enum(SubscriptionStatus, name="subscriptionstatus"),
nullable=False,
default=SubscriptionStatus.EXPIRED,
)
expires_at: Mapped[datetime.datetime] = mapped_column(DateTime(True), nullable=False)
user: Mapped["User"] = relationship("User", back_populates="subscription", lazy="selectin")
addons: Mapped[list["SubscriptionAddon"]] = relationship(
"SubscriptionAddon", back_populates="subscription", lazy="selectin"
)

View File

@@ -1,3 +1,5 @@
import secrets
import string
from typing import TYPE_CHECKING
from sqlalchemy import BIGINT, INTEGER, REAL, TEXT, VARCHAR, ForeignKey
@@ -6,7 +8,12 @@ from sqlalchemy.orm import Mapped, mapped_column, relationship
from db.base import Base
if TYPE_CHECKING:
from db.models.sessions import Session
from db.models import Order, Session, Subscription
def generate_ref_code(length: int = 8):
alphabet = string.ascii_letters + string.digits
return "".join(secrets.choice(alphabet) for _ in range(length))
class User(Base):
@@ -19,11 +26,15 @@ class User(Base):
hashed_password: Mapped[str] = mapped_column(VARCHAR(255), nullable=True)
telegram_id: Mapped[int] = mapped_column(BIGINT, unique=True, nullable=True)
referal_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=True)
referal_code: Mapped[str] = mapped_column(
TEXT, nullable=False, unique=True, default=generate_ref_code
)
balance: Mapped[float] = mapped_column(REAL, nullable=False, default=0)
sessions: Mapped[list["Session"]] = relationship(back_populates="user", lazy="selectin")
referal: Mapped["User | None"] = relationship(
"User",
remote_side=[id],
orders: Mapped[list["Order"]] = relationship("Order", back_populates="user", lazy="selectin")
subscription: Mapped["Subscription"] = relationship(
"Subscription", back_populates="user", lazy="selectin"
)
sessions: Mapped[list["Session"]] = relationship(back_populates="user", lazy="selectin")
referal: Mapped["User | None"] = relationship("User", remote_side=[id], lazy="selectin")

View File

@@ -1,4 +1,4 @@
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from config import cfg
@@ -6,6 +6,29 @@ engine = create_async_engine(cfg.db_url, echo=True)
async_session = async_sessionmaker(bind=engine, expire_on_commit=False)
class UnitOfWork:
def __init__(self, session: AsyncSession) -> None:
self.session = session
async def commit(self) -> None:
await self.session.commit()
async def rollback(self) -> None:
await self.session.rollback()
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, *_):
if exc_type:
await self.rollback()
async def get_db():
async with async_session() as session:
yield session
async def get_uow():
async with async_session() as session, UnitOfWork(session) as uow:
yield uow

View File

@@ -19,8 +19,6 @@ services:
interval: 10s
timeout: 5s
retries: 5
ports:
- "5432:5432"
volumes:
postgres_data:

1
external/pally.py vendored
View File

@@ -226,6 +226,7 @@ class BillService(BaseService):
async def create(
self,
*,
amount: float,
shop_id: str,
order_id: str | None = None,

496
external/rw.py vendored Normal file
View File

@@ -0,0 +1,496 @@
import logging
import re
from dataclasses import dataclass
from datetime import UTC, datetime, timedelta
from uuid import UUID
from remnawave import RemnawaveSDK
from remnawave.models import (
CreateUserRequestDto,
DeleteUserHwidDeviceResponseDto,
HWIDDeleteRequest,
UpdateUserRequestDto,
)
from remnawave.models.hwid import HwidDeviceDto
from config import cfg
logger = logging.getLogger(__name__)
RW_USERNAME_MAX_LENGTH = 36
# ─────────────────────────────────────────────────────────────────
# Dataclass — унифицированное представление пользователя RW
# ─────────────────────────────────────────────────────────────────
@dataclass
class RWUserInfo:
uuid: str
username: str
status: str # "ACTIVE", "DISABLED", "EXPIRED", "ON_HOLD", "LIMITED"
expire_at: datetime
used_traffic_bytes: float
traffic_limit_bytes: int # 0 = unlimited
hwid_device_limit: int | None # None = unlimited
active_squads: list[dict] # [{"uuid": "...", "name": "..."}]
description: str | None
short_uuid: str
# ─────────────────────────────────────────────────────────────────
# Внутренние хелперы
# ─────────────────────────────────────────────────────────────────
def _parse_user(user_dto) -> RWUserInfo:
"""Преобразовать UserResponseDto → RWUserInfo."""
active_squads: list[dict] = []
for squad in user_dto.active_internal_squads or []:
active_squads.append({"uuid": str(squad.uuid), "name": squad.name})
# UserStatus может быть enum-объектом или строкой — нормализуем
status_raw = user_dto.status
if hasattr(status_raw, "value"):
status_str = str(status_raw.value).upper()
else:
status_str = str(status_raw).upper()
# Трафик лежит в user_traffic.used_traffic_bytes
used_traffic: float = 0.0
if user_dto.user_traffic is not None:
used_traffic = float(user_dto.user_traffic.used_traffic_bytes or 0)
return RWUserInfo(
uuid=str(user_dto.uuid),
username=user_dto.username,
status=status_str,
expire_at=user_dto.expire_at,
used_traffic_bytes=used_traffic,
traffic_limit_bytes=int(user_dto.traffic_limit_bytes or 0),
hwid_device_limit=user_dto.hwid_device_limit,
active_squads=active_squads,
description=user_dto.description,
short_uuid=user_dto.short_uuid,
)
def _build_rw_username(*, username: str | None) -> str | None:
if not username:
return None
sanitized_username = re.sub(r"[^a-zA-Z0-9_-]+", "_", username).strip("_")
if not sanitized_username:
return None
return f"web_{sanitized_username}"[:RW_USERNAME_MAX_LENGTH]
# ─────────────────────────────────────────────────────────────────
# Публичные функции — обёртки с перехватом ошибок
# ─────────────────────────────────────────────────────────────────
async def get_user_by_telegram_id(
sdk: RemnawaveSDK,
telegram_id: int,
) -> RWUserInfo | None:
"""
Получить первого пользователя Remnawave, привязанного к Telegram ID.
Возвращает None, если пользователь не найден или Remnawave недоступен.
"""
try:
users_list = await sdk.users.get_users_by_telegram_id(str(telegram_id))
if not users_list:
return None
return _parse_user(users_list[0])
except Exception as exc:
logger.warning(
"Failed to get RW user by telegram_id=%s: %s",
telegram_id,
exc,
)
return None
async def get_user_by_username(
sdk: RemnawaveSDK,
username: str,
) -> RWUserInfo | None:
try:
user = await sdk.users.get_user_by_username(username=username)
if user is None:
return None
return _parse_user(user)
except Exception as exc:
logger.warning(
"Failed to get RW user by username=%s: %s",
username,
exc,
)
return None
async def get_all_squads(sdk: RemnawaveSDK) -> list[dict]:
"""
Получить все Internal Squads из Remnawave.
Возвращает список [{"uuid": "...", "name": "..."}].
При ошибке возвращает пустой список.
"""
try:
response = await sdk.internal_squads.get_internal_squads()
result: list[dict] = []
# GetAllInternalSquadsResponseDto имеет поле .internal_squads
for squad in response.internal_squads:
result.append({"uuid": str(squad.uuid), "name": squad.name})
return result
except Exception as exc:
logger.warning("Failed to get all internal squads: %s", exc)
return []
async def add_days(
sdk: RemnawaveSDK,
user_uuid: str,
days: int,
) -> datetime | None:
"""
Добавить N дней к expire_at пользователя.
Возвращает новую дату окончания или None при ошибке.
"""
try:
# Получаем актуальный expire_at — нельзя слепо использовать
# закешированное значение, чтобы не потерять изменения других агентов
user_dto = await sdk.users.get_user_by_uuid(uuid=user_uuid)
if user_dto is None or user_dto.expire_at is None:
logger.warning("Cannot add days: user %s has no expire_at", user_uuid)
return None
# Обеспечиваем timezone-aware datetime
current_expire = user_dto.expire_at
if current_expire.tzinfo is None:
current_expire = current_expire.replace(tzinfo=UTC)
new_expire = current_expire + timedelta(days=days)
result = await sdk.users.update_user(
UpdateUserRequestDto(
uuid=UUID(user_uuid),
expire_at=new_expire,
)
)
if result is None:
return None
new_dt = result.expire_at
if new_dt is not None and new_dt.tzinfo is None:
new_dt = new_dt.replace(tzinfo=UTC)
return new_dt
except Exception as exc:
logger.warning("Failed to add %d days for user %s: %s", days, user_uuid, exc)
return None
async def remove_days(
sdk: RemnawaveSDK,
user_uuid: str,
days: int,
) -> datetime | None:
"""
Вычесть N дней из expire_at пользователя.
Возвращает новую дату окончания или None при ошибке.
"""
try:
user_dto = await sdk.users.get_user_by_uuid(uuid=user_uuid)
if user_dto is None or user_dto.expire_at is None:
logger.warning("Cannot remove days: user %s has no expire_at", user_uuid)
return None
current_expire = user_dto.expire_at
if current_expire.tzinfo is None:
current_expire = current_expire.replace(tzinfo=UTC)
new_expire = current_expire - timedelta(days=days)
result = await sdk.users.update_user(
UpdateUserRequestDto(
uuid=UUID(user_uuid),
expire_at=new_expire,
)
)
if result is None:
return None
new_dt = result.expire_at
if new_dt is not None and new_dt.tzinfo is None:
new_dt = new_dt.replace(tzinfo=UTC)
return new_dt
except Exception as exc:
logger.warning("Failed to remove %d days for user %s: %s", days, user_uuid, exc)
return None
async def set_hwid_limit(
sdk: RemnawaveSDK,
user_uuid: str,
limit: int,
) -> bool:
"""
Установить лимит HWID-устройств.
limit == 0 → убрать ограничение (None).
Возвращает True при успехе, False при ошибке.
"""
try:
hwid_value: int | None = None if limit == 0 else limit
await sdk.users.update_user(
UpdateUserRequestDto(
uuid=UUID(user_uuid),
hwid_device_limit=hwid_value,
)
)
return True
except Exception as exc:
logger.warning(
"Failed to set HWID limit=%s for user %s: %s",
limit,
user_uuid,
exc,
)
return False
async def set_description(
sdk: RemnawaveSDK,
user_uuid: str,
note: str,
) -> bool:
"""
Установить заметку (description) пользователя в Remnawave.
Возвращает True при успехе, False при ошибке.
"""
try:
await sdk.users.update_user(
UpdateUserRequestDto(
uuid=UUID(user_uuid),
description=note,
)
)
return True
except Exception as exc:
logger.warning(
"Failed to set description for user %s: %s",
user_uuid,
exc,
)
return False
async def reset_traffic(
sdk: RemnawaveSDK,
user_uuid: str,
) -> bool:
"""
Сбросить счётчик использованного трафика пользователя.
Возвращает True при успехе, False при ошибке.
"""
try:
await sdk.users.reset_user_traffic(uuid=user_uuid)
return True
except Exception as exc:
logger.warning("Failed to reset traffic for user %s: %s", user_uuid, exc)
return False
async def disable_user(
sdk: RemnawaveSDK,
user_uuid: str,
) -> bool:
"""
Отключить пользователя (статус → DISABLED).
Возвращает True при успехе, False при ошибке.
"""
try:
await sdk.users.disable_user(uuid=user_uuid)
return True
except Exception as exc:
logger.warning("Failed to disable user %s: %s", user_uuid, exc)
return False
async def enable_user(
sdk: RemnawaveSDK,
user_uuid: str,
) -> bool:
"""
Включить пользователя (статус → ACTIVE).
Возвращает True при успехе, False при ошибке.
"""
try:
await sdk.users.enable_user(uuid=user_uuid)
return True
except Exception as exc:
logger.warning("Failed to enable user %s: %s", user_uuid, exc)
return False
async def update_squads(
sdk: RemnawaveSDK,
user_uuid: str,
squad_uuid_list: list[str],
) -> bool:
"""
Полностью заменить список Internal Squads пользователя.
squad_uuid_list — строковые UUID всех squad'ов (полный новый список).
Возвращает True при успехе, False при ошибке.
"""
try:
await sdk.users.update_user(
UpdateUserRequestDto(
uuid=UUID(user_uuid),
active_internal_squads=[UUID(s) for s in squad_uuid_list],
)
)
return True
except Exception as exc:
logger.warning(
"Failed to update squads for user %s: %s",
user_uuid,
exc,
)
return False
async def create_user(
*,
sdk: RemnawaveSDK,
username: str,
telegram_id: int | None,
expire_at: datetime,
hwid_device_limit: int | None,
squad_uuids: list[str],
) -> RWUserInfo | None:
"""
Создать нового пользователя в Remnawave.
hwid_device_limit=None → без ограничений.
Возвращает RWUserInfo при успехе, None при ошибке.
"""
try:
dto = CreateUserRequestDto(
username=username,
telegram_id=telegram_id,
expire_at=expire_at,
hwid_device_limit=hwid_device_limit,
active_internal_squads=[UUID(s) for s in squad_uuids] if squad_uuids else None,
)
result = await sdk.users.create_user(dto)
if result is None:
return None
return _parse_user(result)
except Exception as exc:
logger.warning("Failed to create user username=%s: %s", username, exc)
return None
async def get_hwid_list(sdk: RemnawaveSDK, user_uuid: str) -> list[HwidDeviceDto] | None:
try:
resp = await sdk.hwid.get_hwid_user(user_uuid)
return resp.devices
except Exception:
logger.exception("failed to fetch hwid list for user=%s", user_uuid)
return
async def delete_hwid(sdk: RemnawaveSDK, user_uuid: str, hwid: str) -> bool:
body = HWIDDeleteRequest(user_uuid=user_uuid, hwid=hwid)
try:
resp = await sdk.hwid.delete_hwid_to_user(body)
return isinstance(resp, DeleteUserHwidDeviceResponseDto)
except Exception:
logger.exception("failed to delete hwid=%s for user=%s", hwid, user_uuid)
return False
async def update_expire_at(sdk: RemnawaveSDK, user_uuid: str, expire_at: datetime):
dto = UpdateUserRequestDto(uuid=UUID(user_uuid), expire_at=expire_at)
try:
await sdk.users.update_user(body=dto)
return True
except Exception:
logger.exception("failed to update expire_at=%s for user=%s", str(expire_at), user_uuid)
return False
def get_sdk() -> RemnawaveSDK | None:
if not cfg.remnawave_base_url or not cfg.remnawave_token:
return None
return RemnawaveSDK(base_url=cfg.remnawave_base_url, token=cfg.remnawave_token)
async def get_rw_user(
sdk: RemnawaveSDK, telegram_id: int | None = None, username: str | None = None
) -> RWUserInfo | None:
rw_user = None
if telegram_id is not None:
rw_user = await get_user_by_telegram_id(sdk, telegram_id)
if rw_user is None:
rw_username = _build_rw_username(username=username)
if rw_username is None:
logger.warning(
"Cannot build username for telegram_id=%s username=%s",
telegram_id,
username,
)
return None
rw_user = await get_user_by_username(sdk, rw_username)
return rw_user
return rw_user
async def sync_subscription_by_telegram_id(
*,
expires_at: datetime,
devices: int,
telegram_id: int | None = None,
username: str | None = None,
) -> bool:
sdk = get_sdk()
if sdk is None:
logger.info("Skipping RW subscription sync: RemnaWave is not configured")
return False
rw_user = await get_rw_user(sdk, telegram_id=telegram_id, username=username)
if expires_at <= datetime.now(UTC):
if rw_user is None:
return True
return await disable_user(sdk, rw_user.uuid)
if rw_user is None:
rw_username = _build_rw_username(username=username)
rw_user = await create_user(
sdk=sdk,
username=rw_username,
telegram_id=telegram_id,
expire_at=expires_at,
hwid_device_limit=None if devices == 0 else devices,
squad_uuids=cfg.remnawave_default_squads,
)
if rw_user is None:
logger.warning(
"Failed to create RW user for telegram_id=%s username=%s",
telegram_id,
username,
)
return False
expire_synced = await update_expire_at(sdk, rw_user.uuid, expires_at)
devices_synced = await set_hwid_limit(sdk, rw_user.uuid, devices)
enabled = await enable_user(sdk, rw_user.uuid)
return expire_synced and devices_synced and enabled
def build_subscription_link(short_uuid: str):
return cfg.remnawave_sub_url.rstrip("/") + f"/{short_uuid}"

25
main.py
View File

@@ -1,8 +1,31 @@
import asyncio
from contextlib import asynccontextmanager, suppress
from fastapi import FastAPI
from routes import routers
from services.notifications import run_subscription_notifications
from services.rw_sync import run_rw_sync_reconciler, run_rw_sync_worker
app = FastAPI(debug=True)
@asynccontextmanager
async def lifespan(_: FastAPI):
tasks = [
asyncio.create_task(run_subscription_notifications()),
asyncio.create_task(run_rw_sync_worker()),
asyncio.create_task(run_rw_sync_reconciler()),
]
try:
yield
finally:
for task in tasks:
task.cancel()
for task in tasks:
with suppress(asyncio.CancelledError):
await task
app = FastAPI(debug=True, lifespan=lifespan)
for r in routers:
app.include_router(r)

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Addon
from db.session import UnitOfWork
class AddonsRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.session = uow.session
async def get_all(self) -> list[Addon]:
stmt = select(Addon)

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.invoice import Invoice
from db.session import UnitOfWork
from schemas.invoices import InvoiceStatus
class InvoiceRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_by_id(self, id: int) -> Invoice | None:
stmt = select(Invoice).where(Invoice.id == id)
@@ -21,21 +22,20 @@ class InvoiceRepository:
return list(r.scalars().all())
async def create(self, creator_id: int, amount: int | float, status: InvoiceStatus) -> Invoice:
async def create(
self, creator_id: int, order_id: int, amount: int | float, status: InvoiceStatus
) -> Invoice:
obj = Invoice(
creator_id=creator_id,
order_id=order_id,
amount=amount,
status=status,
)
self.session.add(obj)
await self.session.commit()
return obj
async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None:
invoice = await self.get_by_id(invoice_id)
invoice.status = status
await self.session.commit()
return invoice

View File

@@ -0,0 +1,38 @@
from datetime import datetime
from sqlalchemy import select
from db.models.link_codes import LinkCode
from db.session import UnitOfWork
from schemas.enums import LinkCodeStatus
async def create_link_code(
uow: UnitOfWork,
*,
code: str,
user_id: int,
status: LinkCodeStatus,
expires_at: datetime,
) -> LinkCode:
link_code = LinkCode(
code=code,
user_id=user_id,
status=status,
expires_at=expires_at,
)
uow.session.add(link_code)
return link_code
async def get_link_code_by_code(uow: UnitOfWork, code: str) -> LinkCode | None:
stmt = select(LinkCode).where(LinkCode.code == code)
r = await uow.session.execute(stmt)
return r.scalar_one_or_none()
async def use_link_code(uow: UnitOfWork, code: LinkCode) -> LinkCode:
code.status = LinkCodeStatus.USED
return code

75
repositories/orders.py Normal file
View File

@@ -0,0 +1,75 @@
from sqlalchemy import select
from db.models.orders import Order, OrderAddon, OrderStatus
from db.session import UnitOfWork
class OrderRepository:
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def create(
self,
*,
user_id: int,
devices: int,
duration_days: int,
total_amount: float,
balance_amount: float,
addons: list[str] | None = None,
status: OrderStatus = OrderStatus.PENDING,
) -> Order:
if not addons:
addons = []
order = Order(
user_id=user_id,
devices=devices,
duration_days=duration_days,
total_amount=total_amount,
balance_amount=balance_amount,
status=status,
)
self.session.add(order)
await self.session.flush()
for addon_id in addons:
addon = OrderAddon(order_id=order.id, addon_id=addon_id)
self.session.add(addon)
await self.session.refresh(order, attribute_names=["addons"])
return order
async def get_by_id(self, order_id: int) -> Order | None:
stmt = select(Order).where(Order.id == order_id)
r = await self.session.execute(stmt)
return r.scalar_one_or_none()
async def get_order_by_user_id(self, user_id: int) -> list[Order]:
stmt = select(Order).where(Order.user_id == user_id)
r = await self.session.execute(stmt)
return list(r.scalars().all())
async def get_active_by_user_id(self, user_id: int) -> list[Order]:
stmt = (
select(Order).where(Order.user_id == user_id).where(Order.status == OrderStatus.PENDING)
)
r = await self.session.execute(stmt)
return list(r.scalars().all())
async def get_paid_unapplied_by_user_id(self, user_id: int) -> list[Order]:
stmt = (
select(Order)
.where(Order.user_id == user_id)
.where(Order.status == OrderStatus.PAID)
.where(Order.applied_at.is_(None))
.order_by(Order.applies_at.asc(), Order.id.asc())
)
r = await self.session.execute(stmt)
return list(r.scalars().all())

View File

@@ -1,12 +1,12 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.pricing import PricingConfig
from db.session import UnitOfWork
class PricingRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.session = uow.session
async def get(self) -> PricingConfig | None:
stmt = select(PricingConfig).where(PricingConfig.id == 1)

View File

@@ -0,0 +1,61 @@
from sqlalchemy import func, or_, select, text
from db.models.service_notifications import ServiceNotification
from db.session import UnitOfWork
from schemas.enums import NotificationStatus
async def get_notification_by_id(uow: UnitOfWork, n_id: int) -> ServiceNotification | None:
stmt = select(ServiceNotification).where(ServiceNotification.id == n_id)
r = await uow.session.execute(stmt)
return r.scalar_one_or_none()
async def get_pending_notifications(
uow: UnitOfWork, batch_size: int = 50
) -> list[ServiceNotification]:
stmt = (
select(ServiceNotification)
.where(
or_(
ServiceNotification.status == "pending",
(
(ServiceNotification.status == "dispatched")
& (
ServiceNotification.dispatched_at
< func.now() - text("interval '10 minutes'")
)
),
)
)
.order_by(ServiceNotification.created_at)
.with_for_update(skip_locked=True)
.limit(batch_size)
)
r = await uow.session.execute(stmt)
return list(r.scalars().all())
async def ack_notification(uow: UnitOfWork, n_id: int) -> ServiceNotification | None:
notification = await get_notification_by_id(uow, n_id)
if not notification:
return
notification.sent_at = func.now()
notification.status = NotificationStatus.SENT
return notification
async def mark_notification_as_dispatched(uow: UnitOfWork, n_id: int):
notification = await get_notification_by_id(uow, n_id)
if not notification:
return
notification.dispatched_at = func.now()
notification.status = NotificationStatus.DISPATCHED
notification.attempts += 1
return notification

View File

@@ -0,0 +1,16 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models.service_signatures import ServiceSignature
from schemas.enums import ServiceSignatureStatus
async def get_active_signature_by_kid(session: AsyncSession, kid: str) -> ServiceSignature | None:
stmt = (
select(ServiceSignature)
.where(ServiceSignature.kid == kid)
.where(ServiceSignature.status == ServiceSignatureStatus.ACTIVE)
)
r = await session.execute(stmt)
return r.scalar_one_or_none()

View File

@@ -1,13 +1,14 @@
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Session
from db.session import UnitOfWork
from schemas.providers import ProvidersType
class SessionsRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_session_by_id(self, id: int) -> Session | None:
stmt = select(Session).where(Session.id == id)
@@ -27,12 +28,10 @@ class SessionsRepository:
async def create(self, user_id: int, refresh_token_hash: str, iss: ProvidersType) -> Session:
obj = Session(user_id=user_id, refresh_token_hash=refresh_token_hash, source=iss)
self.session.add(obj)
await self.session.commit()
return obj
async def revoke(self, token_id: int):
session = await self.get_session_by_id(token_id)
session.is_revoked = True
session.revoked_at = func.now()
await self.session.commit()
return session

View File

@@ -1,13 +1,14 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import User
from db.models.transactions import BalanceTransaction, BalanceTxType
from db.session import UnitOfWork
class UserRepository:
def __init__(self, session: AsyncSession) -> None:
self.session = session
def __init__(self, uow: UnitOfWork) -> None:
self.uow = uow
self.session = uow.session
async def get_user_by_id(self, id: int) -> User | None:
stmt = select(User).where(User.id == id)
@@ -24,21 +25,26 @@ class UserRepository:
res = await self.session.execute(stmt)
return res.scalar_one_or_none()
async def get_user_by_ref_code(self, ref_code: str) -> User | None:
stmt = select(User).where(User.referal_code == ref_code)
res = await self.session.execute(stmt)
return res.scalar_one_or_none()
async def create(
self,
*,
username: str | None = None,
hashed_password: str | None = None,
telegram_id: int | None = None,
referal_id: int | None = None,
) -> User:
obj = User(
username=username,
hashed_password=hashed_password,
telegram_id=telegram_id,
referal_id=referal_id,
)
self.session.add(obj)
await self.session.commit()
return obj
async def increase_balance(
@@ -60,5 +66,8 @@ class UserRepository:
self.session.add(obj)
user.balance += amount
await self.session.commit()
return user
async def update_telegram_id(self, user: User, telegram_id: int) -> User:
user.telegram_id = telegram_id
return user

View File

@@ -7,4 +7,11 @@ pydantic-settings>=2.14.0
asyncpg>=0.31.0
alembic>=1.18.0
aiohttp>=3.14.0
python-multipart==0.0.32
python-multipart==0.0.32
# Private remnawave SDK: bare URL without creds (safe to commit).
# Local install: export REMNAWAVE_SDK_TOKEN=xxx && git config --global url."https://agony:${REMNAWAVE_SDK_TOKEN}@git.mdevs.lat/".insteadOf "https://git.mdevs.lat/" && pip install -r requirements.txt
# Without the token git clone fails with 401/403 (repo is private, login is hardcoded to 'agony').
remnawave @ git+https://git.mdevs.lat/agony/remnawave-sdk.git
zxcvbn>=4.5.0
httpx2>=2.11.0
pytest>=9.1.0

View File

@@ -1,8 +1,21 @@
from fastapi import APIRouter
from .auth import router as auth_router
from .health import router as health_router
from .internal import internal_router
from .link_codes import router as link_code_router
from .orders import router as orders_router
from .payments import payment_routers
from .payments import payment_router
from .plans import router as plans_router
from .users import router as users_routers
routers: list[APIRouter] = [auth_router, plans_router, orders_router, *payment_routers]
routers: list[APIRouter] = [
auth_router,
plans_router,
orders_router,
users_routers,
health_router,
link_code_router,
payment_router,
internal_router,
]

View File

@@ -1,12 +1,15 @@
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from core.secrets import hash_password, hash_refresh_token, verify_password
from db.session import get_db
from core.secrets import (
estimate_password_strength,
hash_password,
hash_refresh_token,
verify_password,
)
from db.session import UnitOfWork, get_uow
from repositories.sessions import SessionsRepository
from repositories.users import UserRepository
from schemas.login import UserLogin, UserLoginData, UserTokens
from schemas.login import AuthenticatedUser, UserLoginData, UserTokens
from schemas.providers import ProvidersType
from schemas.registration import UserRegistration
from schemas.user import UserInfo
@@ -16,9 +19,10 @@ from services.users import authorize_user
router = APIRouter(prefix="/auth")
@router.post("/signup")
async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
@router.post("/signup", response_model=AuthenticatedUser)
async def signup(req: UserRegistration, uow: UnitOfWork = Depends(get_uow)) -> AuthenticatedUser:
users_repo = UserRepository(uow)
sessions_repo = SessionsRepository(uow)
if req.provider == "credentials":
if not req.username or not req.password:
@@ -27,19 +31,40 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
if user:
raise HTTPException(status_code=409, detail="User already exists")
if not estimate_password_strength(req.password):
raise HTTPException(422, detail="Password is not secure.")
password_hash = hash_password(req.password)
user = await users_repo.create(username=req.username, hashed_password=password_hash)
return JSONResponse(
UserInfo(username=user.username, telegram_id=user.telegram_id).model_dump(),
status_code=201,
referal_id = None
if req.referal_code:
referal = await users_repo.get_user_by_ref_code(req.referal_code)
referal_id = referal.id if referal else None
user = await users_repo.create(
username=req.username, hashed_password=password_hash, referal_id=referal_id
)
await uow.commit()
info = UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
bonus_balance=user.balance,
)
user_login = await authorize_user(sessions_repo, user, req.provider)
return AuthenticatedUser(
access_token=user_login.access_token,
refresh_token=user_login.refresh_token,
expires_at=user_login.expires_at,
user=info,
)
raise HTTPException(status_code=400, detail="Unsupported provider")
@router.post("/login", response_model=UserLogin)
async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
users_repo = UserRepository(session)
sessions_repo = SessionsRepository(session)
@router.post("/login", response_model=AuthenticatedUser)
async def login(req: UserLoginData, uow: UnitOfWork = Depends(get_uow)):
users_repo = UserRepository(uow)
sessions_repo = SessionsRepository(uow)
if req.provider == "credentials":
if not req.username or not req.password:
raise HTTPException(status_code=400, detail="Username or password is not provided.")
@@ -52,6 +77,7 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
raise HTTPException(status_code=401, detail="Invalid password")
data = await authorize_user(sessions_repo, user, req.provider)
await uow.commit()
return data
if req.provider == "telegram":
@@ -62,8 +88,8 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
@router.post("/refresh", response_model=UserTokens)
async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession = Depends(get_db)):
sessions_repo = SessionsRepository(session)
async def refresh(refresh_token: str, iss: ProvidersType, uow: UnitOfWork = Depends(get_uow)):
sessions_repo = SessionsRepository(uow)
token_hash = hash_refresh_token(refresh_token)
token_entry = await sessions_repo.get_session_by_hash(token_hash)
@@ -72,6 +98,7 @@ async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession
raise HTTPException(status_code=401, detail="Refresh token is invalid.")
key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss)
await uow.commit()
return UserTokens(
access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token,

8
routes/health.py Normal file
View File

@@ -0,0 +1,8 @@
from fastapi import APIRouter
router = APIRouter(prefix="/health")
@router.get("/")
async def healthcheck():
return "OK"

View File

@@ -0,0 +1,6 @@
from fastapi import APIRouter
from .renewal import router as renewal_router
internal_router = APIRouter(prefix="/internal")
internal_router.include_router(renewal_router)

View File

@@ -0,0 +1,60 @@
from fastapi import APIRouter, Depends, HTTPException
from core.deps import get_service_identity
from db.session import UnitOfWork, get_uow
from repositories.service_notifications import (
ack_notification,
get_pending_notifications,
mark_notification_as_dispatched,
)
from schemas.dto import ServiceIdentity
from schemas.notifications import (
NotificationAcknowledgeRequest,
NotificationResponse,
UserNotificationData,
)
router = APIRouter(prefix="/renewal")
@router.get("/pending", response_model=NotificationResponse)
async def get_pending(
limit: int = 50,
ctx: ServiceIdentity = Depends(get_service_identity),
uow: UnitOfWork = Depends(get_uow),
):
notifications = await get_pending_notifications(uow, limit)
response_users: list[UserNotificationData] = []
for notification in notifications:
response_users.append(
UserNotificationData(
notification_id=notification.id,
username=notification.subscription.user.username,
telegram_id=notification.subscription.user.telegram_id,
expires_at=notification.sub_expires_at,
notification_type=notification.notify_type,
)
)
await mark_notification_as_dispatched(uow, notification.id)
await uow.commit()
return NotificationResponse(users=response_users, issued_by=ctx.service)
@router.post("/ack")
async def acknowledge(
req: NotificationAcknowledgeRequest,
ctx: ServiceIdentity = Depends(get_service_identity),
uow: UnitOfWork = Depends(get_uow),
):
try:
r = await ack_notification(uow, req.notification_id)
except Exception as e:
raise HTTPException(500, detail=str(e)) from None
if r:
await uow.commit()
return "OK"
raise HTTPException(500, detail="No such notification found.")

63
routes/link_codes.py Normal file
View File

@@ -0,0 +1,63 @@
import secrets
from datetime import UTC, datetime, timedelta
from fastapi import APIRouter, Depends, HTTPException
from config import cfg
from core.deps import get_auth_context, get_service_identity
from db.session import UnitOfWork, get_uow
from repositories.link_codes import create_link_code, get_link_code_by_code, use_link_code
from repositories.users import UserRepository
from schemas.dto import AuthContext
from schemas.enums import LinkCodeStatus
from schemas.link_codes import LinkCodeConsume, LinkCodeResponse
from schemas.user import UserInfo
router = APIRouter(prefix="/link-codes")
@router.post("", response_model=LinkCodeResponse, status_code=201)
async def gen_link_code(
ctx: AuthContext = Depends(get_auth_context), uow: UnitOfWork = Depends(get_uow)
):
code = secrets.token_urlsafe(cfg.link_code_length)
exp = datetime.now(UTC) + timedelta(minutes=cfg.link_code_ttl)
link_code = await create_link_code(
uow, code=code, user_id=ctx.user.id, status=LinkCodeStatus.ACTIVE, expires_at=exp
)
await uow.commit()
return LinkCodeResponse(code=link_code.code, expires_at=link_code.expires_at)
@router.post("/consume", response_model=UserInfo)
async def consume_link_code(
payload: LinkCodeConsume,
ctx: AuthContext = Depends(get_service_identity),
uow: UnitOfWork = Depends(get_uow),
):
link_code = await get_link_code_by_code(uow, payload.code)
if not link_code:
raise HTTPException(404, detail="Code not found")
if link_code.status != LinkCodeStatus.ACTIVE:
raise HTTPException(404, detail="Code expired or is invalid.")
users_repo = UserRepository(uow)
user = await users_repo.get_user_by_id(link_code.user_id)
if not user:
raise HTTPException(404, detail="User not found")
user = await users_repo.update_telegram_id(user, payload.telegram_id)
await use_link_code(uow, code=link_code)
await uow.commit()
return UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
bonus_balance=user.balance,
)

View File

@@ -1,41 +1,109 @@
import math
from datetime import UTC, datetime
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.deps import get_auth_context, get_pally_client
from db.session import get_db
from db.models.orders import OrderStatus
from db.session import UnitOfWork, get_uow
from external.pally import PallyClient
from repositories import AddonsRepository, PricingRepository
from repositories.invoices import InvoiceRepository
from repositories.orders import OrderRepository
from schemas.checkout import CheckoutResponse
from schemas.dto import AuthContext
from schemas.invoices import InvoiceResponse, InvoiceStatus
from schemas.invoices import InvoiceStatus
from schemas.plans import OrderDetails
from services.plans import calculate_price, get_pricing_model
from services.subscriptions import (
apply_order_now,
deduct_order_balance,
queue_order_for_later,
should_apply_immediately,
)
router = APIRouter(prefix="/orders")
@router.post("/checkout", response_model=InvoiceResponse, status_code=201)
@router.post("/checkout", response_model=CheckoutResponse, status_code=201)
async def checkout(
order: OrderDetails,
ctx: AuthContext = Depends(get_auth_context),
session: AsyncSession = Depends(get_db),
uow: UnitOfWork = Depends(get_uow),
pally: PallyClient = Depends(get_pally_client),
):
addons_repo = AddonsRepository(session)
pricing_repo = PricingRepository(session)
invoices_repo = InvoiceRepository(session)
addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(uow)
invoices_repo = InvoiceRepository(uow)
orders_repo = OrderRepository(uow)
pricing = await get_pricing_model(addons_repo, pricing_repo)
price = await calculate_price(addons_repo=addons_repo, order=order, pricing=pricing)
price = math.ceil(await calculate_price(addons_repo=addons_repo, order=order, pricing=pricing))
bonus_covered = min(ctx.user.balance, price)
amount_to_pay = round(max(price - bonus_covered, 0), 2)
now = datetime.now(UTC)
invoice = await invoices_repo.create(
creator_id=ctx.user.id, amount=price, status=InvoiceStatus.ACTIVE
order_entry = await orders_repo.create(
user_id=ctx.user.id,
devices=order.devices,
duration_days=order.duration_days,
total_amount=price,
balance_amount=bonus_covered,
addons=order.addons,
)
bill = await pally.bills.create(price, cfg.pally_shop_id, order_id=invoice.id)
if amount_to_pay > 0:
await invoices_repo.create(
creator_id=ctx.user.id,
order_id=order_entry.id,
amount=amount_to_pay,
status=InvoiceStatus.ACTIVE,
)
if not (bill.success and bill.link_page_url):
raise HTTPException(500, detail="Failed to create an invoice.")
bill = await pally.bills.create(
amount=float(amount_to_pay),
shop_id=cfg.pally_shop_id,
order_id=str(order_entry.id),
)
payment_link = bill.link_page_url
if not payment_link:
raise HTTPException(500, detail="failed to create invoice")
else:
await deduct_order_balance(
uow.session,
user=ctx.user,
order=order_entry,
description=f"order {order_entry.id} paid from balance",
)
order_entry.status = OrderStatus.PAID
if should_apply_immediately(
subscription=ctx.user.subscription,
order=order_entry,
pricing=pricing,
now=now,
):
await apply_order_now(
uow.session,
user=ctx.user,
order=order_entry,
pricing=pricing,
now=now,
)
else:
await queue_order_for_later(
order=order_entry, subscription=ctx.user.subscription, now=now
)
await uow.commit()
payment_link = None
return InvoiceResponse(success=True, payment_link=bill.link_page_url, amount=float(price))
if amount_to_pay > 0:
await uow.commit()
return CheckoutResponse(
order_id=str(order_entry.id),
total_amount=price,
bonus_paid=bonus_covered,
amount_to_pay=amount_to_pay,
payment_link=payment_link,
)

View File

@@ -1,3 +1,6 @@
from fastapi import APIRouter
from .pally import router as pally_router
payment_routers = [pally_router]
payment_router = APIRouter(prefix="/payments")
payment_router.include_router(pally_router)

View File

@@ -1,28 +1,24 @@
# ruff: noqa: N803
import hashlib
import hmac
import logging
import math
from fastapi import Depends, Form, HTTPException
from fastapi.routing import APIRouter
from sqlalchemy.ext.asyncio import AsyncSession
from config import cfg
from core.deps import get_db
from db.models.transactions import BalanceTxType
from db.session import UnitOfWork, get_uow
from external.pally import BillStatus
from repositories.invoices import InvoiceRepository
from repositories.users import UserRepository
from schemas.invoices import InvoiceStatus
from services.payments import process_subscription_purchase, validate_pally_signature
router = APIRouter(prefix="/payments/pally")
router = APIRouter(prefix="/pally")
logger = logging.getLogger(__name__)
@router.post("/result")
async def pally_callback( # noqa: PLR0911
async def pally_callback(
*,
InvId: str = Form(...),
OutSum: str = Form(...),
@@ -32,7 +28,6 @@ async def pally_callback( # noqa: PLR0911
CurrencyIn: str = Form(...),
custom: str | None = Form(None),
SignatureValue: str = Form(...),
# Optional fields for additional information
AccountType: str | None = Form(None),
AccountNumber: str | None = Form(None),
BalanceAmount: str | None = Form(None),
@@ -43,12 +38,8 @@ async def pally_callback( # noqa: PLR0911
PayerComment: str | None = Form(None),
ErrorCode: int | None = Form(None),
ErrorMessage: str | None = Form(None),
session: AsyncSession = Depends(get_db),
uow: UnitOfWork = Depends(get_uow),
):
users_repo = UserRepository(session)
invoice_repo = InvoiceRepository(session)
invoice_id_str = InvId
logger.info(
"Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, "
"CurrencyIn: %s, custom: %s, BalanceAmount: %s, SignatureValue: %s",
@@ -72,29 +63,18 @@ async def pally_callback( # noqa: PLR0911
TrsId,
)
# Validate signature
raw_string = f"{OutSum}:{InvId}:{cfg.pally_token}"
expected_signature = hashlib.md5(raw_string.encode("utf-8")).hexdigest().upper()
logger.debug("Signature validation for TrsId %s", TrsId)
if not hmac.compare_digest(SignatureValue, expected_signature):
if not validate_pally_signature(OutSum, InvId, SignatureValue):
logger.critical(
"SECURITY ALERT: Invalid signature for TrsId %s - Expected: %s, Received: %s",
"SECURITY ALERT: Invalid signature for TrsId %s",
TrsId,
expected_signature,
SignatureValue,
)
raise HTTPException(403, detail="Invalid signature.")
# Only process successful payments
if Status != BillStatus.SUCCESS:
logger.info("Bill %s skipped: status=%s", TrsId, Status)
return "OK"
logger.info("Processing successfully paid bill %s", TrsId)
# Validate bill ID (from InvId field, which contains the order_id from bill creation)
invoice_id_str = InvId
if not invoice_id_str or not invoice_id_str.isdigit():
logger.critical(
"Invalid or non-numeric bill ID in InvId field for TrsId %s: '%s'",
@@ -103,120 +83,64 @@ async def pally_callback( # noqa: PLR0911
)
return "OK"
# Find bill in database
invoice = await invoice_repo.get_by_id(int(invoice_id_str))
if not invoice:
logger.critical("Bill %s not found in database for TrsId %s", invoice_id_str, TrsId)
return "OK"
# Check if already processed
if invoice.status != InvoiceStatus.ACTIVE:
logger.warning(
"Bill %s (TrsId: %s) is already processed with status: %s",
invoice.id,
TrsId,
invoice.status,
)
return "OK"
amount = int(float(BalanceAmount)) if BalanceAmount is not None else int(float(OutSum))
try:
# Handle fee scenarios: use BalanceAmount if available (net amount after fees),
# otherwise use OutSum (gross amount paid by customer)
if BalanceAmount is not None:
# Customer pays fees - BalanceAmount is the net amount credited to merchant
credited_amount = int(float(BalanceAmount))
gross_amount = int(float(OutSum))
logger.info(
"Customer-pays-fees payment: bill_id=%s, expected=%s, gross_paid=%s, net_credited=%s",
invoice.id,
invoice.amount,
gross_amount,
credited_amount,
)
# Validate that the net credited amount matches our bill amount
if invoice.amount != credited_amount:
logger.error(
"Net amount mismatch for bill %s (TrsId: %s) - Expected: %s, Net credited: %s, Gross paid: %s",
invoice.id,
TrsId,
invoice.amount,
credited_amount,
gross_amount,
)
return "OK"
amount = credited_amount # Credit the net amount (without fees)
else:
# Standard payment - OutSum should match bill amount exactly
amount = int(float(OutSum))
logger.info(
"Standard payment: bill_id=%s, expected=%s, received=%s",
invoice.id,
invoice.amount,
amount,
)
if invoice.amount != amount:
logger.error(
"Amount mismatch for bill %s (TrsId: %s) - Expected: %s, Received: %s",
invoice.id,
TrsId,
invoice.amount,
amount,
)
return "OK"
logger.info(
"Processing payment: bill_id=%s, user_id=%s, amount=%s",
invoice.id,
invoice.creator_id,
amount,
)
# Credit user balance
await users_repo.increase_balance(
invoice.creator_id,
await process_subscription_purchase(
uow,
invoice_id=int(invoice_id_str),
trs_id=TrsId,
amount=amount,
tx_type=BalanceTxType.DEPOSIT,
description=f"payment via PALLY (TrsId: {TrsId})",
)
# Process referral bonus
user = invoice.creator
referal = user.referal
if referal is not None:
referal_amount = math.floor(amount * (cfg.referal_bonus / 100))
await users_repo.increase_balance(
referal,
referal_amount,
tx_type=BalanceTxType.REFERRAL_BONUS,
description=f"referral reward for user {invoice.creator_id} (TrsId: {TrsId})",
)
logger.info(
"Referral bonus processed: referrer_id=%s, amount=%s", referal, referal_amount
)
logger.info("Payment processing completed successfully for TrsId %s", TrsId)
except Exception as e:
except Exception:
# The payment is confirmed, so never leave the user without the money
# if order/subscription provisioning fails.
logger.exception(
"CRITICAL ERROR processing payment for TrsId %s, bill_id %s, user_id %s: %s",
"Subscription provisioning failed for bill %s (TrsId: %s); "
"crediting the user's bonus balance",
invoice_id_str,
TrsId,
invoice.id,
invoice.creator_id,
str(e),
)
# Don't return early - still mark as success to prevent retries
# The balance operation might have partially succeeded
await uow.rollback()
# Update bill status to success
await invoice_repo.update_status_by_id(int(invoice.id), status=InvoiceStatus.PAID)
invoice_repo = InvoiceRepository(uow)
invoice = await invoice_repo.get_by_id(int(invoice_id_str))
if invoice is None:
logger.critical(
"Cannot credit fallback balance: bill %s was not found (TrsId: %s)",
invoice_id_str,
TrsId,
)
raise
logger.info("Bill %s marked as SUCCESS for TrsId %s", invoice.id, TrsId)
if invoice.status == InvoiceStatus.PAID:
return "OK"
users_repo = UserRepository(uow)
user = await users_repo.get_user_by_id(invoice.creator_id)
if user is None:
logger.critical(
"Cannot credit fallback balance: user %s was not found " "for bill %s (TrsId: %s)",
invoice.creator_id,
invoice_id_str,
TrsId,
)
raise
await users_repo.increase_balance(
user.id,
amount,
BalanceTxType.DEPOSIT,
f"fallback payment credit for invoice {invoice.id} (TrsId: {TrsId})",
)
await invoice_repo.update_status_by_id(invoice.id, InvoiceStatus.PAID)
await uow.commit()
logger.info(
"Fallback payment credit processed: user_id=%s, amount=%s, " "invoice_id=%s, TrsId=%s",
user.id,
amount,
invoice.id,
TrsId,
)
return "OK"

View File

@@ -1,7 +1,6 @@
from fastapi import APIRouter, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from db.session import get_db
from db.session import UnitOfWork, get_uow
from repositories.addons import AddonsRepository
from repositories.pricing import PricingRepository
from schemas.plans import PricingPlans
@@ -11,9 +10,9 @@ router = APIRouter(prefix="/plans")
@router.get("/", response_model=PricingPlans)
async def get_plans(session: AsyncSession = Depends(get_db)):
addons_repo = AddonsRepository(session)
pricing_repo = PricingRepository(session)
async def get_plans(uow: UnitOfWork = Depends(get_uow)):
addons_repo = AddonsRepository(uow)
pricing_repo = PricingRepository(uow)
res = await get_pricing_model(addons_repo, pricing_repo)
return res

97
routes/users.py Normal file
View File

@@ -0,0 +1,97 @@
import logging
from fastapi import APIRouter, Depends, HTTPException
from core.deps import get_auth_context
from external.rw import build_subscription_link, delete_hwid, get_hwid_list, get_rw_user, get_sdk
from schemas.common import OperationData
from schemas.devices import Device
from schemas.dto import AuthContext
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(
username=ctx.user.username,
telegram_id=ctx.user.telegram_id,
referal_code=ctx.user.referal_code,
bonus_balance=ctx.user.balance,
)
@router.get("/subscription")
async def get_subscription(ctx: AuthContext = Depends(get_auth_context)):
sub = ctx.user.subscription
if not sub:
return SubscriptionData(
has_subscription=False,
devices=None,
expires_at=None,
addon_ids=[],
subscription_link=None,
duration_days=None,
)
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,
)
return SubscriptionData(
has_subscription=True,
devices=sub.devices,
expires_at=sub.expires_at,
addon_ids=[a.addon_id for a in sub.addons],
subscription_link=None,
duration_days=sub.duration_days,
)
return SubscriptionData(
has_subscription=True,
devices=sub.devices,
expires_at=sub.expires_at,
addon_ids=[a.addon_id for a in sub.addons],
subscription_link=build_subscription_link(rw_user.short_uuid),
duration_days=sub.duration_days,
)
@router.get("/subscription/hwid", response_model=list[Device])
async def get_hwid(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:
raise HTTPException(403, detail="RW user not found.")
hwid_list = await get_hwid_list(get_sdk(), user_uuid=rw_user.uuid)
devices = []
if not hwid_list:
return devices
for device in hwid_list:
devices.append(
Device(
os=device.platform,
model=device.device_model,
client=device.user_agent.split("/")[0],
hwid=device.hwid,
)
)
return devices
@router.delete("/subscription/hwid")
async def delete_hwid_endpoint(hwid: str, 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:
raise HTTPException(403, detail="RW user not found.")
success = await delete_hwid(get_sdk(), rw_user.uuid, hwid)
return OperationData(success=success)

15
schemas/checkout.py Normal file
View File

@@ -0,0 +1,15 @@
from pydantic import BaseModel, Field, computed_field
class CheckoutResponse(BaseModel):
order_id: str = Field()
total_amount: float = Field()
bonus_paid: float = Field()
amount_to_pay: float = Field()
payment_link: str | None = Field(None)
@property
@computed_field
def is_fully_paid(self):
return self.total_amount <= self.bonus_paid and self.amount_to_pay <= 0

5
schemas/common.py Normal file
View File

@@ -0,0 +1,5 @@
from pydantic import BaseModel, Field
class OperationData(BaseModel):
success: bool = Field(False)

8
schemas/devices.py Normal file
View File

@@ -0,0 +1,8 @@
from pydantic import BaseModel, Field
class Device(BaseModel):
os: str | None = Field(None)
model: str | None = Field(None)
client: str | None = Field(None)
hwid: str = Field()

View File

@@ -14,5 +14,14 @@ class KeyPair:
@dataclass
class AuthContext:
user: User
auth_method: Literal["jwt"]
auth_method: Literal["jwt", "service"]
service: str | None = None
@dataclass
class ServiceIdentity:
"""
Proves that request is coming from verified source, no user context provided.
"""
service: str

31
schemas/enums.py Normal file
View File

@@ -0,0 +1,31 @@
from enum import StrEnum
class SubscriptionStatus(StrEnum):
ACTIVE = "active"
EXPIRED = "expired"
class ServiceSignatureStatus(StrEnum):
ACTIVE = "active"
INACTIVE = "inactive"
class LinkCodeStatus(StrEnum):
ACTIVE = "active"
USED = "used"
EXPIRED = "expired"
class NotificationType(StrEnum):
SEVEN_DAYS = "7d"
THREE_DAYS = "3d"
ONE_DAY = "1d"
EXPIRED = "expired"
class NotificationStatus(StrEnum):
PENDING = "pending"
DISPATCHED = "dispatched"
SENT = "sent"
FAILED = "failed"

View File

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

20
schemas/link_codes.py Normal file
View File

@@ -0,0 +1,20 @@
from datetime import datetime
from pydantic import BaseModel, Field, computed_field
from config import cfg
class LinkCodeResponse(BaseModel):
code: str = Field()
expires_at: datetime = Field()
@computed_field
@property
def deep_link(self) -> str:
return cfg.bot_url + "?start=" + self.code
class LinkCodeConsume(BaseModel):
code: str = Field()
telegram_id: int = Field()

View File

@@ -16,7 +16,7 @@ class UserLoginData(BaseModel):
telegram: TelegramData | None = None
class UserLogin(BaseModel):
class AuthenticatedUser(BaseModel):
access_token: str
refresh_token: str
expires_at: float

24
schemas/notifications.py Normal file
View File

@@ -0,0 +1,24 @@
from datetime import UTC, datetime
from pydantic import BaseModel, Field
from schemas.enums import NotificationType
class UserNotificationData(BaseModel):
notification_id: int = Field()
username: str = Field()
telegram_id: int | None = Field()
expires_at: datetime = Field()
notification_type: NotificationType = Field()
class NotificationResponse(BaseModel):
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
issued_by: str = Field()
users: list[UserNotificationData] = Field()
class NotificationAcknowledgeRequest(BaseModel):
notification_id: int = Field()

View File

@@ -6,6 +6,7 @@ class AddonData(BaseModel):
name: str
price: float
free_threshold: int
is_enabled: bool
class PricingPlans(BaseModel):

View File

@@ -9,4 +9,6 @@ class UserRegistration(BaseModel):
username: str | None = Field(None)
password: str | None = Field(None)
referal_code: str | None = Field(None)
provider: ProvidersType

View File

@@ -1,6 +1,19 @@
from datetime import datetime
from pydantic import BaseModel, Field
class SubscriptionData(BaseModel):
has_subscription: bool = Field()
devices: int | None = Field(None)
expires_at: datetime | None = Field(None)
addon_ids: list[str] = Field(default_factory=list)
subscription_link: str | None = Field(None)
duration_days: int | None = Field(None)
class UserInfo(BaseModel):
username: str | None = Field()
telegram_id: str | None = Field()
username: str | None = Field(None)
telegram_id: int | None = Field(None)
referal_code: str = Field()
bonus_balance: float = Field()

59
services/notifications.py Normal file
View File

@@ -0,0 +1,59 @@
import asyncio
from datetime import UTC, datetime, timedelta
from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert
from config import cfg
from db.models import ServiceNotification, Subscription
from db.session import UnitOfWork, async_session
from schemas.enums import NotificationType
async def collect_subscription_notifications(uow: UnitOfWork, now: datetime | None = None) -> int:
session = uow.session
now = now or datetime.now(UTC)
periods = (
(NotificationType.SEVEN_DAYS, now + timedelta(days=3), now + timedelta(days=7)),
(NotificationType.THREE_DAYS, now + timedelta(days=1), now + timedelta(days=3)),
(NotificationType.ONE_DAY, now, now + timedelta(days=1)),
(NotificationType.EXPIRED, None, now),
)
created = 0
for notification_type, starts_at, ends_at in periods:
conditions = [Subscription.expires_at <= ends_at]
if starts_at is not None:
conditions.append(Subscription.expires_at > starts_at)
subscriptions = await session.scalars(select(Subscription).where(*conditions))
rows = [
{
"subscription_id": subscription.id,
"notify_type": notification_type,
"sub_expires_at": subscription.expires_at,
}
for subscription in subscriptions
]
if not rows:
continue
stmt = insert(ServiceNotification).values(rows)
stmt = stmt.on_conflict_do_nothing(constraint="uq_subscription_notification")
result = await session.execute(stmt)
created += result.rowcount or 0
await uow.commit()
return created
async def run_subscription_notifications() -> None:
interval = cfg.notification_scan_interval * 60
while True:
try:
async with async_session() as session, UnitOfWork(session) as uow:
await collect_subscription_notifications(uow)
except Exception:
# A failed iteration must not stop subsequent notification checks.
pass
await asyncio.sleep(interval)

171
services/payments.py Normal file
View File

@@ -0,0 +1,171 @@
import hashlib
import hmac
import logging
import math
from datetime import UTC, datetime
from config import cfg
from db.models.orders import OrderStatus
from db.models.transactions import BalanceTxType
from db.session import UnitOfWork
from repositories import AddonsRepository
from repositories.invoices import InvoiceRepository
from repositories.orders import OrderRepository
from repositories.pricing import PricingRepository
from repositories.users import UserRepository
from schemas.invoices import InvoiceStatus
from services.plans import get_pricing_model
from services.rw_sync import enqueue_rw_sync
from services.subscriptions import (
apply_order_now,
deduct_order_balance,
queue_order_for_later,
should_apply_immediately,
)
logger = logging.getLogger(__name__)
def validate_pally_signature(out_sum: str, inv_id: str, signature_value: str) -> bool:
raw_string = f"{out_sum}:{inv_id}:{cfg.pally_token}"
expected_signature = hashlib.md5(raw_string.encode("utf-8")).hexdigest().upper()
return hmac.compare_digest(signature_value, expected_signature)
async def process_subscription_purchase( # noqa: PLR0911, PLR0912
uow: UnitOfWork,
*,
invoice_id: int,
trs_id: str,
amount: int,
) -> None:
session = uow.session
invoice_repo = InvoiceRepository(uow)
orders_repo = OrderRepository(uow)
users_repo = UserRepository(uow)
pricing_repo = PricingRepository(uow)
invoice = await invoice_repo.get_by_id(invoice_id)
if not invoice:
logger.critical("Bill %s not found in database for TrsId %s", invoice_id, trs_id)
return
if invoice.status != InvoiceStatus.ACTIVE:
logger.warning(
"Bill %s (TrsId: %s) is already processed with status: %s",
invoice.id,
trs_id,
invoice.status,
)
return
if invoice.amount != amount:
logger.error(
"Amount mismatch for bill %s (TrsId: %s) - Expected: %s, Received: %s",
invoice.id,
trs_id,
invoice.amount,
amount,
)
return
if invoice.order_id is None:
logger.critical("Invoice %s has no linked order for TrsId %s", invoice.id, trs_id)
return
order = await orders_repo.get_by_id(invoice.order_id)
if order is None:
logger.critical(
"Order not found for bill %s (TrsId: %s, user_id=%s, order_id=%s)",
invoice.id,
trs_id,
invoice.creator_id,
invoice.order_id,
)
return
if order.user_id != invoice.creator_id:
logger.critical(
"Order %s does not belong to invoice creator %s for TrsId %s",
order.id,
invoice.creator_id,
trs_id,
)
return
user = await users_repo.get_user_by_id(invoice.creator_id)
if user is None:
logger.critical("User %s not found for TrsId %s", invoice.creator_id, trs_id)
return
subscription = user.subscription
now = datetime.now(UTC)
pricing = await get_pricing_model(AddonsRepository(uow), pricing_repo)
logger.info(
"Processing payment: bill_id=%s, order_id=%s, user_id=%s, amount=%s",
invoice.id,
order.id,
invoice.creator_id,
amount,
)
order.status = OrderStatus.PAID
invoice.status = InvoiceStatus.PAID
if order.balance_amount > 0:
await deduct_order_balance(
session,
user=user,
order=order,
description=f"order {order.id} partial payment from balance",
)
if should_apply_immediately(
subscription=subscription,
order=order,
pricing=pricing,
now=now,
):
applied_subscription = await apply_order_now(
session, user=user, order=order, pricing=pricing, now=now
)
await enqueue_rw_sync(session, applied_subscription.id)
else:
if subscription is None:
logger.critical(
"Cannot queue order %s without subscription for user %s", order.id, user.id
)
return
await queue_order_for_later(order=order, subscription=subscription, now=now)
referal_id = user.referal_id
if referal_id is not None:
referal_amount = math.floor(amount * (cfg.referal_bonus / 100))
if referal_amount > 0:
referal_user = await users_repo.get_user_by_id(referal_id)
if referal_user is not None:
await users_repo.increase_balance(
referal_id,
referal_amount,
BalanceTxType.REFERRAL_BONUS,
f"referral reward for user {invoice.creator_id} (TrsId: {trs_id})",
)
logger.info(
"Referral bonus processed: referrer_id=%s, amount=%s",
referal_id,
referal_amount,
)
else:
logger.warning(
"Referrer %s not found for user_id=%s while processing TrsId %s",
referal_id,
invoice.creator_id,
trs_id,
)
await uow.commit()
logger.info("Bill %s marked as PAID for TrsId %s", invoice_id, trs_id)

View File

@@ -1,6 +1,7 @@
from repositories.addons import AddonsRepository
from repositories.pricing import PricingRepository
from schemas.plans import AddonData, OrderDetails, PricingPlans
from services.subscriptions import calculate_order_total
async def get_pricing_model(addons_repo: AddonsRepository, pricing_repo: PricingRepository):
@@ -10,7 +11,13 @@ async def get_pricing_model(addons_repo: AddonsRepository, pricing_repo: Pricing
addons = await addons_repo.get_all()
addons_data = [
AddonData(id=a.id, name=a.name, price=a.price, free_threshold=a.free_threshold)
AddonData(
id=a.id,
name=a.name,
price=a.price,
free_threshold=a.free_threshold,
is_enabled=a.is_enabled,
)
for a in addons
]
@@ -21,4 +28,9 @@ async def calculate_price(
*, addons_repo: AddonsRepository, order: OrderDetails, pricing: PricingPlans
) -> float:
addons = [await addons_repo.get_by_id(a) for a in order.addons]
return order.devices * pricing.device_price + sum([a.price for a in addons])
return calculate_order_total(
pricing,
order.devices,
[addon.id for addon in addons],
order.duration_days,
)

126
services/rw_sync.py Normal file
View File

@@ -0,0 +1,126 @@
import asyncio
import logging
from datetime import UTC, datetime, timedelta
from typing import TYPE_CHECKING
from sqlalchemy import delete, select, update
from sqlalchemy.dialects.postgresql import insert
from config import cfg
from db.models import RWSyncOutbox, Subscription
from db.session import UnitOfWork, async_session
from external.rw import sync_subscription_by_telegram_id
if TYPE_CHECKING:
from sqlalchemy.ext.asyncio import AsyncSession
logger = logging.getLogger(__name__)
LOCK_DURATION = timedelta(minutes=5)
MAX_RETRY_DELAY = timedelta(hours=1)
async def enqueue_rw_sync(session: "AsyncSession", subscription_id: int) -> None:
now = datetime.now(UTC)
stmt = insert(RWSyncOutbox).values(
subscription_id=subscription_id,
revision=1,
next_attempt_at=now,
)
await session.execute(
stmt.on_conflict_do_update(
index_elements=[RWSyncOutbox.subscription_id],
set_={
"revision": RWSyncOutbox.revision + 1,
"next_attempt_at": now,
"locked_until": None,
},
)
)
async def enqueue_all_rw_syncs(session: "AsyncSession") -> int:
subscription_ids = await session.scalars(select(Subscription.id))
count = 0
for subscription_id in subscription_ids:
await enqueue_rw_sync(session, subscription_id)
count += 1
return count
async def process_rw_syncs(uow: UnitOfWork, batch_size: int = 50) -> int:
session = uow.session
now = datetime.now(UTC)
jobs = await session.scalars(
select(RWSyncOutbox)
.where(
RWSyncOutbox.next_attempt_at <= now,
(RWSyncOutbox.locked_until.is_(None)) | (RWSyncOutbox.locked_until <= now),
)
.order_by(RWSyncOutbox.next_attempt_at)
.with_for_update(skip_locked=True)
.limit(batch_size)
)
jobs = list(jobs)
for job in jobs:
job.locked_until = now + LOCK_DURATION
await uow.commit()
for job in jobs:
subscription = job.subscription
user = subscription.user
synced = await sync_subscription_by_telegram_id(
expires_at=subscription.expires_at,
devices=subscription.devices,
telegram_id=user.telegram_id,
username=user.username,
)
if synced:
await session.execute(
delete(RWSyncOutbox).where(
RWSyncOutbox.id == job.id,
RWSyncOutbox.revision == job.revision,
)
)
await uow.commit()
continue
delay = min(timedelta(minutes=2**job.attempts), MAX_RETRY_DELAY)
await session.execute(
update(RWSyncOutbox)
.where(RWSyncOutbox.id == job.id, RWSyncOutbox.revision == job.revision)
.values(
attempts=RWSyncOutbox.attempts + 1,
next_attempt_at=datetime.now(UTC) + delay,
locked_until=None,
)
)
await uow.commit()
logger.warning("RW sync failed for subscription_id=%s", subscription.id)
return len(jobs)
async def run_rw_sync_worker() -> None:
interval = cfg.rw_sync_interval_minutes * 60
while True:
try:
async with async_session() as session, UnitOfWork(session) as uow:
await process_rw_syncs(uow)
except Exception:
logger.exception("RW sync worker iteration failed")
await asyncio.sleep(interval)
async def run_rw_sync_reconciler() -> None:
interval = cfg.rw_sync_reconcile_interval_minutes * 60
while True:
try:
async with async_session() as session:
async with UnitOfWork(session) as uow:
count = await enqueue_all_rw_syncs(uow.session)
await uow.commit()
logger.info("Queued %s subscriptions for RW reconciliation", count)
except Exception:
logger.exception("RW reconciliation iteration failed")
await asyncio.sleep(interval)

197
services/subscriptions.py Normal file
View File

@@ -0,0 +1,197 @@
from datetime import UTC, datetime, timedelta
from sqlalchemy import delete
from sqlalchemy.ext.asyncio import AsyncSession
from db.models import Subscription, SubscriptionAddon, User
from db.models.orders import Order, OrderStatus
from db.models.transactions import BalanceTransaction, BalanceTxType
from schemas.enums import SubscriptionStatus
from schemas.plans import PricingPlans
from services.rw_sync import enqueue_rw_sync
def calculate_plan_monthly_price(
pricing: PricingPlans, devices: int, addon_ids: list[str]
) -> float:
addon_prices = {addon.id: addon.price for addon in pricing.addons}
return devices * pricing.device_price + sum(addon_prices[addon_id] for addon_id in addon_ids)
def calculate_order_total(
pricing: PricingPlans, devices: int, addon_ids: list[str], duration_days: int
) -> float:
monthly_price = calculate_plan_monthly_price(pricing, devices, addon_ids)
return monthly_price * duration_days / 30
def get_subscription_addon_ids(subscription: Subscription | None) -> list[str]:
if subscription is None:
return []
return [addon.addon_id for addon in subscription.addons]
def should_apply_immediately(
*,
order: Order,
pricing: PricingPlans,
now: datetime,
subscription: Subscription | None = None,
) -> bool:
if subscription is None or subscription.expires_at <= now:
return True
current_addons = set(get_subscription_addon_ids(subscription))
new_addons = {addon.addon_id for addon in order.addons}
if order.devices == subscription.devices and new_addons == current_addons:
return True
current_monthly_price = calculate_plan_monthly_price(
pricing, subscription.devices, list(current_addons)
)
new_monthly_price = calculate_plan_monthly_price(pricing, order.devices, list(new_addons))
return (
order.devices >= subscription.devices
and new_addons.issuperset(current_addons)
and new_monthly_price >= current_monthly_price
)
async def replace_subscription_addons(
session: AsyncSession, subscription_id: int, addon_ids: list[str]
) -> None:
await session.execute(
delete(SubscriptionAddon).where(SubscriptionAddon.subscription_id == subscription_id)
)
for addon_id in addon_ids:
session.add(SubscriptionAddon(subscription_id=subscription_id, addon_id=addon_id))
async def ensure_subscription(
user: User, session: AsyncSession, starts_at: datetime, duration_days: int
) -> Subscription:
subscription = user.subscription
if subscription is not None:
return subscription
subscription = Subscription(
user_id=user.id,
devices=0,
status=SubscriptionStatus.EXPIRED,
expires_at=starts_at,
duration_days=duration_days,
)
session.add(subscription)
await session.flush()
user.subscription = subscription
return subscription
async def deduct_order_balance(
session: AsyncSession,
*,
user: User,
order: Order,
description: str,
) -> None:
if order.balance_amount <= 0:
return
session.add(
BalanceTransaction(
user_id=user.id,
amount=-order.balance_amount,
tx_type=BalanceTxType.PURCHASE,
balance_before=user.balance,
balance_after=user.balance - order.balance_amount,
description=description,
)
)
user.balance -= order.balance_amount
async def apply_order_now(
session: AsyncSession,
*,
user: User,
order: Order,
pricing: PricingPlans,
now: datetime,
) -> Subscription:
subscription = await ensure_subscription(user, session, now, order.duration_days)
addon_ids = [addon.addon_id for addon in order.addons]
current_addons = [] if subscription.devices == 0 else get_subscription_addon_ids(subscription)
if subscription.status == SubscriptionStatus.ACTIVE and subscription.expires_at > now:
if order.devices == subscription.devices and set(addon_ids) == set(current_addons):
subscription.expires_at += timedelta(days=order.duration_days)
else:
current_monthly_price = calculate_plan_monthly_price(
pricing, subscription.devices, current_addons
)
next_monthly_price = calculate_plan_monthly_price(pricing, order.devices, addon_ids)
remaining_seconds = (subscription.expires_at - now).total_seconds()
remaining_days = max(remaining_seconds / 86400, 0)
remaining_credit = current_monthly_price * remaining_days / 30
purchased_value = calculate_order_total(
pricing, order.devices, addon_ids, order.duration_days
)
total_days = ((remaining_credit + purchased_value) / next_monthly_price) * 30
subscription.expires_at = now + timedelta(days=total_days)
else:
subscription.expires_at = now + timedelta(days=order.duration_days)
subscription.devices = order.devices
subscription.status = SubscriptionStatus.ACTIVE
await replace_subscription_addons(session, subscription.id, addon_ids)
order.applies_at = now
order.applied_at = now
return subscription
async def queue_order_for_later(*, order: Order, subscription: Subscription, now: datetime) -> None:
order.applies_at = max(subscription.expires_at, now)
async def sync_user_subscription(
session: AsyncSession,
*,
user: User,
) -> None:
now = datetime.now(UTC)
subscription = user.subscription
applied_due_orders = False
if subscription is not None and subscription.expires_at <= now:
due_orders = [
order
for order in user.orders
if order.status == OrderStatus.PAID
and order.applied_at is None
and order.applies_at is not None
and order.applies_at <= now
]
due_orders.sort(key=lambda order: (order.applies_at, order.id))
if due_orders:
for order in due_orders:
base_time = max(subscription.expires_at, order.applies_at)
subscription.devices = order.devices
subscription.status = SubscriptionStatus.ACTIVE
subscription.expires_at = base_time + timedelta(days=order.duration_days)
await replace_subscription_addons(
session,
subscription.id,
[addon.addon_id for addon in order.addons],
)
order.applied_at = now
applied_due_orders = True
else:
subscription.status = SubscriptionStatus.EXPIRED
if applied_due_orders:
await enqueue_rw_sync(session, subscription.id)

View File

@@ -1,23 +1,28 @@
from core.secrets import generate_pair, hash_refresh_token
from db.models.users import User
from repositories.sessions import SessionsRepository
from schemas.login import UserLogin
from schemas.login import AuthenticatedUser
from schemas.providers import ProvidersType
from schemas.user import UserInfo
async def authorize_user(
sessions_repo: SessionsRepository, user: User, iss: ProvidersType
) -> UserLogin:
) -> AuthenticatedUser:
key_pair = generate_pair(user.id, iss)
refresh_token_hash = hash_refresh_token(key_pair.refresh_token)
await sessions_repo.create(user_id=user.id, refresh_token_hash=refresh_token_hash, iss=iss)
return UserLogin(
return AuthenticatedUser(
access_token=key_pair.access_token,
refresh_token=key_pair.refresh_token,
user=UserInfo(username=user.username, telegram_id=user.telegram_id),
user=UserInfo(
username=user.username,
telegram_id=user.telegram_id,
referal_code=user.referal_code,
bonus_balance=user.balance,
),
expires_at=key_pair.expires_at,
)

17
tests/conftest.py Normal file
View File

@@ -0,0 +1,17 @@
import sys
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from main import app
@pytest.fixture
def client():
app.dependency_overrides.clear()
with TestClient(app, raise_server_exceptions=False) as test_client:
yield test_client
app.dependency_overrides.clear()

127
tests/test_auth_routes.py Normal file
View File

@@ -0,0 +1,127 @@
# ruff: noqa: PLR2004
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
def test_signup_rejects_missing_credentials(client):
response = client.post("/auth/signup", json={"provider": "credentials"})
assert response.status_code == 400
assert response.json()["detail"] == "Username or password is not provided"
def test_signup_rejects_unsupported_provider(client):
response = client.post(
"/auth/signup", json={"provider": "telegram", "username": "alice", "password": "Strong123!"}
)
assert response.status_code == 400
assert response.json()["detail"] == "Unsupported provider"
def test_signup_rejects_existing_username(client):
repository = SimpleNamespace(get_user_by_username=AsyncMock(return_value=object()))
with patch("routes.auth.UserRepository", return_value=repository):
response = client.post(
"/auth/signup",
json={"provider": "credentials", "username": "alice", "password": "Strong123!"},
)
assert response.status_code == 409
assert response.json()["detail"] == "User already exists"
def test_signup_rejects_weak_password(client):
repository = SimpleNamespace(get_user_by_username=AsyncMock(return_value=None))
with (
patch("routes.auth.UserRepository", return_value=repository),
patch("routes.auth.estimate_password_strength", return_value=False),
):
response = client.post(
"/auth/signup",
json={"provider": "credentials", "username": "alice", "password": "weak"},
)
assert response.status_code == 422
assert response.json()["detail"] == "Password is not secure."
def test_signup_accepts_unknown_referral_code(client):
created_user = SimpleNamespace(
id=1, username="alice", telegram_id=None, referal_code="new-code", balance=0
)
repository = SimpleNamespace(
get_user_by_username=AsyncMock(return_value=None),
get_user_by_ref_code=AsyncMock(return_value=None),
create=AsyncMock(return_value=created_user),
)
sessions_repository = SimpleNamespace(create=AsyncMock(return_value=None))
with (
patch("routes.auth.UserRepository", return_value=repository),
patch("routes.auth.SessionsRepository", return_value=sessions_repository),
patch("routes.auth.estimate_password_strength", return_value=True),
patch("routes.auth.hash_password", return_value="hashed"),
):
response = client.post(
"/auth/signup",
json={
"provider": "credentials",
"username": "alice",
"password": "Strong123!",
"referal_code": "unknown",
},
)
assert response.status_code == 200
repository.create.assert_awaited_once_with(
username="alice", hashed_password="hashed", referal_id=None
)
assert response.json()["user"]["referal_code"] == "new-code"
def test_login_distinguishes_unknown_user_and_bad_password(client):
unknown_repository = SimpleNamespace(get_user_by_username=AsyncMock(return_value=None))
with patch("routes.auth.UserRepository", return_value=unknown_repository):
unknown_response = client.post(
"/auth/login",
json={"provider": "credentials", "username": "alice", "password": "secret"},
)
existing_user = SimpleNamespace(hashed_password="hash")
existing_repository = SimpleNamespace(
get_user_by_username=AsyncMock(return_value=existing_user)
)
with (
patch("routes.auth.UserRepository", return_value=existing_repository),
patch("routes.auth.SessionsRepository"),
patch("routes.auth.verify_password", return_value=False),
):
password_response = client.post(
"/auth/login",
json={"provider": "credentials", "username": "alice", "password": "secret"},
)
assert unknown_response.status_code == password_response.status_code == 401
assert unknown_response.json()["detail"] == "User doesn't exist."
assert password_response.json()["detail"] == "Invalid password"
def test_login_telegram_is_explicitly_unavailable(client):
response = client.post("/auth/login", json={"provider": "telegram"})
assert response.status_code == 503
def test_refresh_requires_query_parameters_and_rejects_unknown_token(client):
missing_response = client.post("/auth/refresh")
repository = SimpleNamespace(get_session_by_hash=AsyncMock(return_value=None))
with (
patch("routes.auth.SessionsRepository", return_value=repository),
patch("routes.auth.hash_refresh_token", return_value="token-hash"),
):
invalid_response = client.post("/auth/refresh?refresh_token=expired&iss=credentials")
assert missing_response.status_code == 422
assert invalid_response.status_code == 401
assert invalid_response.json()["detail"] == "Refresh token is invalid."

View File

@@ -0,0 +1,137 @@
# ruff: noqa: PLR2004
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from config import cfg
from routes import link_codes
from schemas.enums import LinkCodeStatus
def auth_context(user_id=42):
return SimpleNamespace(user=SimpleNamespace(id=user_id))
def service_identity():
return SimpleNamespace(service="telegram-bot")
def test_generate_link_code_requires_authorization(client):
response = client.post("/link-codes")
assert response.status_code == 403
assert response.json()["detail"] == "No authorization provided."
def test_generate_link_code_creates_active_code_with_expiry_and_deep_link(client):
client.app.dependency_overrides[link_codes.get_auth_context] = auth_context
created_codes = []
async def create_code(session, *, code, user_id, status, expires_at):
created_codes.append(
SimpleNamespace(
session=session,
code=code,
user_id=user_id,
status=status,
expires_at=expires_at,
)
)
return created_codes[-1]
before = datetime.now(UTC)
with (
patch(
"routes.link_codes.secrets.token_urlsafe", return_value="one-time-code"
) as token_urlsafe,
patch("routes.link_codes.create_link_code", side_effect=create_code),
):
response = client.post("/link-codes")
after = datetime.now(UTC)
assert response.status_code == 201
assert response.json() == {
"code": "one-time-code",
"expires_at": created_codes[0].expires_at.isoformat().replace("+00:00", "Z"),
"deep_link": f"{cfg.bot_url}?start=one-time-code",
}
assert created_codes[0].user_id == 42
assert created_codes[0].status is LinkCodeStatus.ACTIVE
token_urlsafe.assert_called_once_with(cfg.link_code_length)
assert before + timedelta(minutes=cfg.link_code_ttl) <= created_codes[0].expires_at
assert created_codes[0].expires_at <= after + timedelta(minutes=cfg.link_code_ttl)
def test_consume_link_code_requires_code_and_telegram_id(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
missing_code = client.post("/link-codes/consume", json={"telegram_id": 12345})
missing_telegram_id = client.post("/link-codes/consume", json={"code": "link-code"})
invalid_telegram_id = client.post(
"/link-codes/consume", json={"code": "link-code", "telegram_id": "not-an-id"}
)
assert missing_code.status_code == 422
assert missing_telegram_id.status_code == 422
assert invalid_telegram_id.status_code == 422
def test_consume_link_code_rejects_unknown_code(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
lookup = AsyncMock(return_value=None)
with patch("routes.link_codes.get_link_code_by_code", lookup):
response = client.post(
"/link-codes/consume", json={"code": "missing", "telegram_id": 12345}
)
assert response.status_code == 404
assert response.json()["detail"] == "Code not found"
lookup.assert_awaited_once()
assert lookup.await_args.args[1] == "missing"
def test_consume_link_code_rejects_code_for_deleted_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE)
repository = SimpleNamespace(get_user_by_id=AsyncMock(return_value=None))
with (
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)),
patch("routes.link_codes.UserRepository", return_value=repository),
):
response = client.post(
"/link-codes/consume", json={"code": "orphaned", "telegram_id": 12345}
)
assert response.status_code == 404
assert response.json()["detail"] == "User not found"
repository.get_user_by_id.assert_awaited_once_with(42)
def test_consume_link_code_updates_telegram_id_and_returns_user(client):
client.app.dependency_overrides[link_codes.get_service_identity] = service_identity
link_code = SimpleNamespace(user_id=42, status=LinkCodeStatus.ACTIVE)
user = SimpleNamespace(
username="alice", telegram_id=12345, referal_code="ref-code", balance=100
)
repository = SimpleNamespace(
get_user_by_id=AsyncMock(return_value=user),
update_telegram_id=AsyncMock(return_value=user),
)
with (
patch("routes.link_codes.get_link_code_by_code", new=AsyncMock(return_value=link_code)),
patch("routes.link_codes.UserRepository", return_value=repository),
):
response = client.post(
"/link-codes/consume", json={"code": "link-code", "telegram_id": 12345}
)
assert response.status_code == 200
assert response.json() == {
"username": "alice",
"telegram_id": 12345,
"referal_code": "ref-code",
"bonus_balance": 100,
}
repository.get_user_by_id.assert_awaited_once_with(42)
repository.update_telegram_id.assert_awaited_once_with(user, 12345)

View File

@@ -0,0 +1,47 @@
import asyncio
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from schemas.enums import NotificationType
from services.notifications import collect_subscription_notifications
class FakeUnitOfWork:
def __init__(self, subscriptions_by_period):
self.session = self
self.subscriptions_by_period = iter(subscriptions_by_period)
self.executed = []
self.committed = False
async def scalars(self, _statement):
return next(self.subscriptions_by_period)
async def execute(self, statement):
self.executed.append(statement)
return SimpleNamespace(rowcount=1)
async def commit(self):
self.committed = True
def test_collects_notification_for_each_expiry_period():
now = datetime(2026, 8, 20, tzinfo=UTC)
uow = FakeUnitOfWork(
[
[SimpleNamespace(id=1, expires_at=now + timedelta(days=6))],
[SimpleNamespace(id=2, expires_at=now + timedelta(days=2))],
[SimpleNamespace(id=3, expires_at=now + timedelta(hours=12))],
[SimpleNamespace(id=4, expires_at=now - timedelta(hours=1))],
]
)
created = asyncio.run(collect_subscription_notifications(uow, now))
assert created == len(uow.executed)
assert uow.committed
assert [statement.compile().params["notify_type_m0"] for statement in uow.executed] == [
NotificationType.SEVEN_DAYS,
NotificationType.THREE_DAYS,
NotificationType.ONE_DAY,
NotificationType.EXPIRED,
]

View File

@@ -0,0 +1,106 @@
# ruff: noqa: PLR2004
import hashlib
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from external.pally import BillStatus
from schemas.plans import AddonData, PricingPlans
from services.payments import validate_pally_signature
from services.subscriptions import calculate_order_total, should_apply_immediately
def pricing():
return PricingPlans(
device_price=100,
addons=[
AddonData(id="a", name="A", price=25, free_threshold=0, is_enabled=True),
AddonData(id="b", name="B", price=50, free_threshold=0, is_enabled=True),
],
)
def test_price_calculation_handles_fractional_months_and_unknown_addons():
assert calculate_order_total(pricing(), devices=2, addon_ids=["a"], duration_days=15) == 112.5
with pytest.raises(KeyError):
calculate_order_total(pricing(), devices=1, addon_ids=["missing"], duration_days=30)
def test_should_apply_immediately_distinguishes_upgrade_from_downgrade():
now = datetime.now(UTC)
subscription = SimpleNamespace(
devices=2,
expires_at=now + timedelta(days=10),
addons=[SimpleNamespace(addon_id="a")],
)
upgrade = SimpleNamespace(
devices=3, addons=[SimpleNamespace(addon_id="a"), SimpleNamespace(addon_id="b")]
)
downgrade = SimpleNamespace(devices=1, addons=[SimpleNamespace(addon_id="a")])
assert should_apply_immediately(
order=upgrade, subscription=subscription, pricing=pricing(), now=now
)
assert not should_apply_immediately(
order=downgrade, subscription=subscription, pricing=pricing(), now=now
)
assert should_apply_immediately(order=downgrade, subscription=None, pricing=pricing(), now=now)
def test_pally_signature_is_exact_and_case_sensitive():
signature = hashlib.md5(b"10:42:test-token").hexdigest().upper()
with patch("services.payments.cfg.pally_token", "test-token"):
assert validate_pally_signature("10", "42", signature)
assert not validate_pally_signature("10", "42", signature.lower())
assert not validate_pally_signature("11", "42", signature)
def callback_data(**overrides):
data = {
"InvId": "42",
"OutSum": "10",
"Commission": "0",
"TrsId": "transaction",
"Status": BillStatus.SUCCESS,
"CurrencyIn": "RUB",
"SignatureValue": "valid",
}
data.update(overrides)
return data
def test_callback_rejects_invalid_signature_before_processing(client):
with patch("routes.payments.pally.validate_pally_signature", return_value=False):
response = client.post("/payments/pally/result", data=callback_data())
assert response.status_code == 403
assert response.json()["detail"] == "Invalid signature."
def test_callback_acknowledges_valid_non_success_and_non_numeric_invoice(client):
with patch("routes.payments.pally.validate_pally_signature", return_value=True):
failed_response = client.post("/payments/pally/result", data=callback_data(Status="Failed"))
invalid_id_response = client.post(
"/payments/pally/result", data=callback_data(InvId="order-42")
)
assert failed_response.status_code == invalid_id_response.status_code == 200
assert failed_response.json() == invalid_id_response.json() == "OK"
def test_callback_uses_balance_amount_when_present(client):
process = AsyncMock()
with (
patch("routes.payments.pally.validate_pally_signature", return_value=True),
patch("routes.payments.pally.process_subscription_purchase", process),
):
response = client.post("/payments/pally/result", data=callback_data(BalanceAmount="7.9"))
assert response.status_code == 200
process.assert_awaited_once()
assert process.await_args.kwargs["invoice_id"] == 42
assert process.await_args.kwargs["amount"] == 7

View File

@@ -0,0 +1,131 @@
# ruff: noqa: PLR2004, PLW0108
from datetime import UTC, datetime
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from routes import users
def auth_context(subscription=None):
return SimpleNamespace(
user=SimpleNamespace(
id=1,
username="alice",
telegram_id=12345,
referal_code="ref-code",
balance=12.5,
subscription=subscription,
)
)
def test_protected_user_endpoint_requires_authorization(client):
response = client.get("/users/me")
assert response.status_code == 403
def test_get_me_serializes_auth_context(client):
client.app.dependency_overrides[users.get_auth_context] = lambda: auth_context()
response = client.get("/users/me")
assert response.status_code == 200
assert response.json() == {
"username": "alice",
"telegram_id": 12345,
"referal_code": "ref-code",
"bonus_balance": 12.5,
}
def test_subscription_without_local_subscription_has_empty_details(client):
client.app.dependency_overrides[users.get_auth_context] = lambda: auth_context()
response = client.get("/users/subscription")
assert response.status_code == 200
assert response.json() == {
"has_subscription": False,
"devices": None,
"expires_at": None,
"addon_ids": [],
"subscription_link": None,
"duration_days": None,
}
def test_subscription_keeps_local_data_when_external_user_is_missing(client):
subscription = SimpleNamespace(
id=2,
devices=3,
expires_at=datetime(2030, 1, 1, tzinfo=UTC),
duration_days=30,
addons=[SimpleNamespace(addon_id="a"), SimpleNamespace(addon_id="b")],
)
client.app.dependency_overrides[users.get_auth_context] = lambda: auth_context(subscription)
with (
patch("routes.users.get_sdk", return_value="sdk"),
patch("routes.users.get_rw_user", new=AsyncMock(return_value=None)),
):
response = client.get("/users/subscription")
assert response.status_code == 200
assert response.json()["has_subscription"] is True
assert response.json()["addon_ids"] == ["a", "b"]
assert response.json()["subscription_link"] is None
def test_hwid_returns_empty_list_and_rejects_missing_external_user(client):
client.app.dependency_overrides[users.get_auth_context] = lambda: auth_context()
with (
patch("routes.users.get_sdk", return_value="sdk"),
patch("routes.users.get_rw_user", new=AsyncMock(return_value=None)),
):
missing_user = client.get("/users/subscription/hwid")
rw_user = SimpleNamespace(uuid="uuid")
with (
patch("routes.users.get_sdk", return_value="sdk"),
patch("routes.users.get_rw_user", new=AsyncMock(return_value=rw_user)),
patch("routes.users.get_hwid_list", new=AsyncMock(return_value=None)),
):
empty_list = client.get("/users/subscription/hwid")
assert missing_user.status_code == 403
assert empty_list.status_code == 200
assert empty_list.json() == []
def test_hwid_maps_client_prefix_and_delete_preserves_external_result(client):
client.app.dependency_overrides[users.get_auth_context] = lambda: auth_context()
rw_user = SimpleNamespace(uuid="uuid")
device = SimpleNamespace(
platform="Windows", device_model="PC", user_agent="v2ray/6.0", hwid="abc"
)
with (
patch("routes.users.get_sdk", return_value="sdk"),
patch("routes.users.get_rw_user", new=AsyncMock(return_value=rw_user)),
patch("routes.users.get_hwid_list", new=AsyncMock(return_value=[device])),
patch("routes.users.delete_hwid", new=AsyncMock(return_value=False)),
):
devices_response = client.get("/users/subscription/hwid")
delete_response = client.delete("/users/subscription/hwid?hwid=abc")
missing_hwid_response = client.delete("/users/subscription/hwid")
assert devices_response.json() == [
{"os": "Windows", "model": "PC", "client": "v2ray", "hwid": "abc"}
]
assert delete_response.status_code == 200
assert delete_response.json() == {"success": False}
assert missing_hwid_response.status_code == 422
def test_plans_returns_pricing_service_result(client):
expected = {"device_price": 100, "addons": []}
with patch("routes.plans.get_pricing_model", new=AsyncMock(return_value=expected)):
response = client.get("/plans/")
assert response.status_code == 200
assert response.json() == expected