Compare commits
36 Commits
cde8ce28f5
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 4262d4d017 | |||
| 429c961ace | |||
| 27e2f58956 | |||
| a374501e4d | |||
| 8d1b753b99 | |||
| b39cae8046 | |||
| 0a58a41930 | |||
| 05ded7d7ea | |||
| e212492b30 | |||
| b759a997d9 | |||
| 77ef4aaa57 | |||
| 2e01b6502c | |||
| 771d44d34c | |||
| 7ae3c98585 | |||
| 938e924107 | |||
| 3b7606107b | |||
| 3d79ffb384 | |||
| 039babf540 | |||
| 9e209dd695 | |||
| 05914f0b23 | |||
| 8b6c43f89a | |||
| 865c5cb9b3 | |||
| 9fa8e8c9a9 | |||
| 9a62722a3c | |||
| 5d610c484d | |||
| 8cb1ac89a6 | |||
| d1f7612f37 | |||
| 03e6363c1a | |||
| c06be7b5f2 | |||
| 801014c392 | |||
| 680c15414c | |||
| dee2cad540 | |||
| 74e952698b | |||
| 1e4cb43ac7 | |||
| 97a4e819d6 | |||
| 24f857ef1b |
1
.gitignore
vendored
1
.gitignore
vendored
@@ -180,4 +180,3 @@ plans
|
|||||||
dev.sh
|
dev.sh
|
||||||
test_dummy.py
|
test_dummy.py
|
||||||
*.pem
|
*.pem
|
||||||
tests/
|
|
||||||
1
.python-version
Normal file
1
.python-version
Normal file
@@ -0,0 +1 @@
|
|||||||
|
3.13.0
|
||||||
46
alembic/versions/0ea5625b4913_order_durationdays.py
Normal file
46
alembic/versions/0ea5625b4913_order_durationdays.py
Normal 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 ###
|
||||||
34
alembic/versions/135672cdf14a_user_referal_code.py
Normal file
34
alembic/versions/135672cdf14a_user_referal_code.py
Normal 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 ###
|
||||||
46
alembic/versions/426b372286c3_add_rw_sync_outbox.py
Normal file
46
alembic/versions/426b372286c3_add_rw_sync_outbox.py
Normal 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 ###
|
||||||
41
alembic/versions/551c0ad261cd_service_signatures.py
Normal file
41
alembic/versions/551c0ad261cd_service_signatures.py
Normal file
@@ -0,0 +1,41 @@
|
|||||||
|
"""+service_signatures
|
||||||
|
|
||||||
|
Revision ID: 551c0ad261cd
|
||||||
|
Revises: cc6625f7dd7f
|
||||||
|
Create Date: 2026-08-18 19:17:33.861149
|
||||||
|
|
||||||
|
"""
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = '551c0ad261cd'
|
||||||
|
down_revision: Union[str, Sequence[str], None] = 'cc6625f7dd7f'
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
"""Upgrade schema."""
|
||||||
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
|
op.create_table('service_signatures',
|
||||||
|
sa.Column('kid', sa.TEXT(), nullable=False),
|
||||||
|
sa.Column('public_key', sa.TEXT(), nullable=False),
|
||||||
|
sa.Column('service_name', sa.TEXT(), nullable=True),
|
||||||
|
sa.Column('status', sa.Enum('ACTIVE', 'INACTIVE', name='servicesignaturestatus'), nullable=False),
|
||||||
|
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||||
|
sa.Column('revoked_at', sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.PrimaryKeyConstraint('kid'),
|
||||||
|
sa.UniqueConstraint('kid')
|
||||||
|
)
|
||||||
|
# ### end Alembic commands ###
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
"""Downgrade schema."""
|
||||||
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
|
op.drop_table('service_signatures')
|
||||||
|
# ### end Alembic commands ###
|
||||||
44
alembic/versions/6d1bae3dc723_link_codes.py
Normal file
44
alembic/versions/6d1bae3dc723_link_codes.py
Normal 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 ###
|
||||||
@@ -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 ###
|
||||||
@@ -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 ###
|
||||||
@@ -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"],
|
||||||
|
)
|
||||||
72
alembic/versions/b97ce5b7d663_subscriptions_infra.py
Normal file
72
alembic/versions/b97ce5b7d663_subscriptions_infra.py
Normal 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 ###
|
||||||
32
alembic/versions/bc424721d767_order_user_id_not_unique.py
Normal file
32
alembic/versions/bc424721d767_order_user_id_not_unique.py
Normal 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 ###
|
||||||
@@ -22,7 +22,7 @@ def upgrade() -> None:
|
|||||||
"""Upgrade schema."""
|
"""Upgrade schema."""
|
||||||
# ### commands auto generated by Alembic - please adjust! ###
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
op.create_unique_constraint(None, 'invoices', ['id'])
|
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 ###
|
# ### end Alembic commands ###
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
32
alembic/versions/cc6625f7dd7f_addon_is_enabled.py
Normal file
32
alembic/versions/cc6625f7dd7f_addon_is_enabled.py
Normal 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 ###
|
||||||
47
alembic/versions/ee6174eeef14_service_notifications.py
Normal file
47
alembic/versions/ee6174eeef14_service_notifications.py
Normal 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 ###
|
||||||
32
alembic/versions/f4ece5936e3d_subscription_duration_days.py
Normal file
32
alembic/versions/f4ece5936e3d_subscription_duration_days.py
Normal 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
4
compose.local.yml
Normal file
@@ -0,0 +1,4 @@
|
|||||||
|
services:
|
||||||
|
postgres:
|
||||||
|
ports: !override
|
||||||
|
- "5432:5432"
|
||||||
40
config.py
40
config.py
@@ -5,21 +5,50 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
|
|||||||
|
|
||||||
|
|
||||||
class Settings(BaseSettings):
|
class Settings(BaseSettings):
|
||||||
model_config = SettingsConfigDict(env_file=".env")
|
model_config = SettingsConfigDict(env_file=".env", extra="ignore")
|
||||||
|
|
||||||
|
### Internal settings ###
|
||||||
postgres_user: str = Field()
|
postgres_user: str = Field()
|
||||||
postgres_password: str = Field()
|
postgres_password: str = Field()
|
||||||
postgres_host: str = Field()
|
postgres_host: str = Field()
|
||||||
postgres_port: str = Field()
|
postgres_port: str = Field()
|
||||||
postgres_db: 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()
|
private_key_fp: str = Field()
|
||||||
public_key_fp: str = Field()
|
public_key_fp: str = Field()
|
||||||
|
|
||||||
access_token_ttl: int = Field(description="Access token TTL (minutes)")
|
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()
|
min_password_length: int = Field(8)
|
||||||
pally_token: str = Field()
|
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
|
@computed_field
|
||||||
@property
|
@property
|
||||||
@@ -42,5 +71,10 @@ class Settings(BaseSettings):
|
|||||||
with open(self.public_key_fp, "rb") as f:
|
with open(self.public_key_fp, "rb") as f:
|
||||||
return f.read()
|
return f.read()
|
||||||
|
|
||||||
|
@computed_field
|
||||||
|
@property
|
||||||
|
def bot_url(self) -> str:
|
||||||
|
return "https://t.me/" + self.bot_username.lstrip("@")
|
||||||
|
|
||||||
|
|
||||||
cfg = Settings() # type: ignore
|
cfg = Settings() # type: ignore
|
||||||
|
|||||||
14
core/auth/fetch_sub.py
Normal file
14
core/auth/fetch_sub.py
Normal 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
|
||||||
@@ -2,26 +2,56 @@ from datetime import UTC, datetime
|
|||||||
|
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException
|
||||||
from pydantic import ValidationError
|
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 repositories.users import UserRepository
|
||||||
from schemas.dto import AuthContext
|
from schemas.dto import AuthContext
|
||||||
from schemas.jwt import JWTPayload
|
from schemas.jwt import ServiceJWTPayload, UserJWTPayload
|
||||||
|
|
||||||
|
|
||||||
async def authorize(token: str, session: AsyncSession, service: str | None = None) -> AuthContext:
|
async def authorize_bot(kid: str, token: str, uow: UnitOfWork) -> AuthContext:
|
||||||
content = decode_jwt(token)
|
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:
|
try:
|
||||||
payload = JWTPayload.model_validate(content)
|
payload = ServiceJWTPayload.model_validate(content)
|
||||||
except ValidationError:
|
except ValidationError:
|
||||||
raise HTTPException(status_code=401, detail="Invalid credentials") from None
|
raise HTTPException(status_code=401, detail="Invalid credentials") from None
|
||||||
|
|
||||||
if payload.exp < datetime.now(UTC).timestamp():
|
if payload.exp < datetime.now(UTC).timestamp():
|
||||||
raise HTTPException(status_code=401, detail="Access token expired")
|
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))
|
user = await repo.get_user_by_id(int(payload.sub))
|
||||||
|
|
||||||
if not user:
|
if not user:
|
||||||
|
|||||||
33
core/deps.py
33
core/deps.py
@@ -1,15 +1,17 @@
|
|||||||
from fastapi import Depends, HTTPException, Request
|
from fastapi import Depends, HTTPException, Request
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from config import cfg
|
from config import cfg
|
||||||
from core.auth import jwt
|
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 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(
|
async def get_auth_context(
|
||||||
request: Request, session: AsyncSession = Depends(get_db)
|
request: Request, uow: UnitOfWork = Depends(get_uow)
|
||||||
) -> AuthContext | None:
|
) -> AuthContext | None:
|
||||||
auth = request.headers.get("Authorization")
|
auth = request.headers.get("Authorization")
|
||||||
|
|
||||||
@@ -18,7 +20,28 @@ async def get_auth_context(
|
|||||||
|
|
||||||
if auth.startswith("Bearer"):
|
if auth.startswith("Bearer"):
|
||||||
token = auth.removeprefix("Bearer ").strip()
|
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:
|
def get_pally_client() -> PallyClient:
|
||||||
|
|||||||
@@ -1,15 +1,16 @@
|
|||||||
import hashlib
|
import hashlib
|
||||||
import logging
|
import logging
|
||||||
import secrets
|
import secrets
|
||||||
from typing import Any
|
from typing import Any, Literal
|
||||||
|
|
||||||
import jwt
|
import jwt
|
||||||
from argon2 import PasswordHasher
|
from argon2 import PasswordHasher
|
||||||
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
||||||
|
from zxcvbn import zxcvbn
|
||||||
|
|
||||||
from config import cfg
|
from config import cfg
|
||||||
from schemas.dto import KeyPair
|
from schemas.dto import KeyPair
|
||||||
from schemas.jwt import JWTPayload
|
from schemas.jwt import UserJWTPayload
|
||||||
from schemas.providers import ProvidersType
|
from schemas.providers import ProvidersType
|
||||||
|
|
||||||
ctx = PasswordHasher()
|
ctx = PasswordHasher()
|
||||||
@@ -35,15 +36,29 @@ def generate_jwt(payload: dict[str, Any]) -> str:
|
|||||||
return jwt.encode(payload, cfg.private_key, "EdDSA")
|
return jwt.encode(payload, cfg.private_key, "EdDSA")
|
||||||
|
|
||||||
|
|
||||||
def decode_jwt(token: str) -> dict[str, Any] | None:
|
def get_kid_from_token(token: str) -> str | None:
|
||||||
try:
|
try:
|
||||||
return jwt.decode(token, cfg.public_key, "EdDSA")
|
header = jwt.get_unverified_header(token)
|
||||||
except jwt.ExpiredSignatureError:
|
return header.get("kid")
|
||||||
|
except jwt.exceptions.PyJWTError:
|
||||||
return
|
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:
|
def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair:
|
||||||
payload = JWTPayload(sub=str(user_id), iss=iss)
|
payload = UserJWTPayload(sub=str(user_id), iss=iss)
|
||||||
access_token = generate_jwt(payload.model_dump())
|
access_token = generate_jwt(payload.model_dump())
|
||||||
refresh_token = secrets.token_urlsafe(32)
|
refresh_token = secrets.token_urlsafe(32)
|
||||||
|
|
||||||
@@ -52,3 +67,11 @@ def generate_pair(user_id: int, iss: ProvidersType) -> KeyPair:
|
|||||||
|
|
||||||
def hash_refresh_token(token: str):
|
def hash_refresh_token(token: str):
|
||||||
return hashlib.sha256(token.encode()).hexdigest()
|
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
|
||||||
|
|||||||
@@ -1,8 +1,30 @@
|
|||||||
from .addons import Addon
|
from .addons import Addon
|
||||||
from .invoice import Invoice
|
from .invoice import Invoice
|
||||||
|
from .link_codes import LinkCode
|
||||||
|
from .orders import Order, OrderAddon
|
||||||
from .pricing import PricingConfig
|
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 .sessions import Session
|
||||||
|
from .subscription_addons import SubscriptionAddon
|
||||||
|
from .subscriptions import Subscription
|
||||||
from .transactions import BalanceTransaction
|
from .transactions import BalanceTransaction
|
||||||
from .users import User
|
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",
|
||||||
|
]
|
||||||
|
|||||||
@@ -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 sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
from db.base import Base
|
from db.base import Base
|
||||||
@@ -11,3 +11,4 @@ class Addon(Base):
|
|||||||
name: Mapped[str] = mapped_column(TEXT, nullable=False, unique=False)
|
name: Mapped[str] = mapped_column(TEXT, nullable=False, unique=False)
|
||||||
price: Mapped[float] = mapped_column(FLOAT, nullable=False)
|
price: Mapped[float] = mapped_column(FLOAT, nullable=False)
|
||||||
free_threshold: Mapped[int] = mapped_column(INTEGER, default=-1)
|
free_threshold: Mapped[int] = mapped_column(INTEGER, default=-1)
|
||||||
|
is_enabled: Mapped[bool] = mapped_column(BOOLEAN, default=False, nullable=False)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from db.base import Base
|
|||||||
from schemas.invoices import InvoiceStatus
|
from schemas.invoices import InvoiceStatus
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from db.models import User
|
from db.models import Order, User
|
||||||
|
|
||||||
|
|
||||||
class Invoice(Base):
|
class Invoice(Base):
|
||||||
@@ -17,9 +17,11 @@ class Invoice(Base):
|
|||||||
INTEGER, autoincrement=True, unique=True, nullable=False, primary_key=True
|
INTEGER, autoincrement=True, unique=True, nullable=False, primary_key=True
|
||||||
)
|
)
|
||||||
creator_id: Mapped[int] = mapped_column(ForeignKey("users.id"), nullable=False)
|
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)
|
amount: Mapped[float] = mapped_column(FLOAT, nullable=False)
|
||||||
status: Mapped[InvoiceStatus] = mapped_column(
|
status: Mapped[InvoiceStatus] = mapped_column(
|
||||||
Enum(InvoiceStatus, name="invoicestatus"), nullable=False, default=InvoiceStatus.ACTIVE
|
Enum(InvoiceStatus, name="invoicestatus"), nullable=False, default=InvoiceStatus.ACTIVE
|
||||||
)
|
)
|
||||||
|
|
||||||
creator: Mapped["User"] = relationship("User", lazy="selectin")
|
creator: Mapped["User"] = relationship("User", lazy="selectin")
|
||||||
|
order: Mapped["Order | None"] = relationship("Order", lazy="selectin")
|
||||||
|
|||||||
28
db/models/link_codes.py
Normal file
28
db/models/link_codes.py
Normal 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
53
db/models/orders.py
Normal 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"
|
||||||
|
)
|
||||||
34
db/models/rw_sync_outbox.py
Normal file
34
db/models/rw_sync_outbox.py
Normal 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")
|
||||||
52
db/models/service_notifications.py
Normal file
52
db/models/service_notifications.py
Normal 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")
|
||||||
26
db/models/service_signatures.py
Normal file
26
db/models/service_signatures.py
Normal file
@@ -0,0 +1,26 @@
|
|||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from sqlalchemy import TEXT, DateTime, Enum
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from db.base import Base
|
||||||
|
from schemas.enums import ServiceSignatureStatus
|
||||||
|
|
||||||
|
|
||||||
|
class ServiceSignature(Base):
|
||||||
|
__tablename__ = "service_signatures"
|
||||||
|
|
||||||
|
kid: Mapped[str] = mapped_column(TEXT, unique=True, nullable=False, primary_key=True)
|
||||||
|
public_key: Mapped[str] = mapped_column(TEXT, nullable=False)
|
||||||
|
service_name: Mapped[str] = mapped_column(TEXT, nullable=True)
|
||||||
|
status: Mapped[ServiceSignatureStatus] = mapped_column(
|
||||||
|
Enum(ServiceSignatureStatus, name="servicesignaturestatus"),
|
||||||
|
nullable=False,
|
||||||
|
default=ServiceSignatureStatus.INACTIVE,
|
||||||
|
)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
nullable=False,
|
||||||
|
default=lambda: datetime.now(UTC),
|
||||||
|
)
|
||||||
|
revoked_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||||
23
db/models/subscription_addons.py
Normal file
23
db/models/subscription_addons.py
Normal 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"
|
||||||
|
)
|
||||||
34
db/models/subscriptions.py
Normal file
34
db/models/subscriptions.py
Normal 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"
|
||||||
|
)
|
||||||
@@ -1,3 +1,5 @@
|
|||||||
|
import secrets
|
||||||
|
import string
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from sqlalchemy import BIGINT, INTEGER, REAL, TEXT, VARCHAR, ForeignKey
|
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
|
from db.base import Base
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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):
|
class User(Base):
|
||||||
@@ -19,11 +26,15 @@ class User(Base):
|
|||||||
hashed_password: Mapped[str] = mapped_column(VARCHAR(255), nullable=True)
|
hashed_password: Mapped[str] = mapped_column(VARCHAR(255), nullable=True)
|
||||||
telegram_id: Mapped[int] = mapped_column(BIGINT, unique=True, 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_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)
|
balance: Mapped[float] = mapped_column(REAL, nullable=False, default=0)
|
||||||
|
|
||||||
sessions: Mapped[list["Session"]] = relationship(back_populates="user", lazy="selectin")
|
orders: Mapped[list["Order"]] = relationship("Order", back_populates="user", lazy="selectin")
|
||||||
referal: Mapped["User | None"] = relationship(
|
subscription: Mapped["Subscription"] = relationship(
|
||||||
"User",
|
"Subscription", back_populates="user", lazy="selectin"
|
||||||
remote_side=[id],
|
|
||||||
)
|
)
|
||||||
|
sessions: Mapped[list["Session"]] = relationship(back_populates="user", lazy="selectin")
|
||||||
|
referal: Mapped["User | None"] = relationship("User", remote_side=[id], lazy="selectin")
|
||||||
|
|||||||
@@ -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
|
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)
|
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 def get_db():
|
||||||
async with async_session() as session:
|
async with async_session() as session:
|
||||||
yield session
|
yield session
|
||||||
|
|
||||||
|
|
||||||
|
async def get_uow():
|
||||||
|
async with async_session() as session, UnitOfWork(session) as uow:
|
||||||
|
yield uow
|
||||||
|
|||||||
@@ -19,8 +19,6 @@ services:
|
|||||||
interval: 10s
|
interval: 10s
|
||||||
timeout: 5s
|
timeout: 5s
|
||||||
retries: 5
|
retries: 5
|
||||||
ports:
|
|
||||||
- "5432:5432"
|
|
||||||
|
|
||||||
volumes:
|
volumes:
|
||||||
postgres_data:
|
postgres_data:
|
||||||
|
|||||||
1
external/pally.py
vendored
1
external/pally.py
vendored
@@ -226,6 +226,7 @@ class BillService(BaseService):
|
|||||||
|
|
||||||
async def create(
|
async def create(
|
||||||
self,
|
self,
|
||||||
|
*,
|
||||||
amount: float,
|
amount: float,
|
||||||
shop_id: str,
|
shop_id: str,
|
||||||
order_id: str | None = None,
|
order_id: str | None = None,
|
||||||
|
|||||||
496
external/rw.py
vendored
Normal file
496
external/rw.py
vendored
Normal 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
25
main.py
@@ -1,8 +1,31 @@
|
|||||||
|
import asyncio
|
||||||
|
from contextlib import asynccontextmanager, suppress
|
||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
|
||||||
from routes import routers
|
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:
|
for r in routers:
|
||||||
app.include_router(r)
|
app.include_router(r)
|
||||||
|
|||||||
@@ -1,12 +1,12 @@
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from db.models import Addon
|
from db.models import Addon
|
||||||
|
from db.session import UnitOfWork
|
||||||
|
|
||||||
|
|
||||||
class AddonsRepository:
|
class AddonsRepository:
|
||||||
def __init__(self, session: AsyncSession) -> None:
|
def __init__(self, uow: UnitOfWork) -> None:
|
||||||
self.session = session
|
self.session = uow.session
|
||||||
|
|
||||||
async def get_all(self) -> list[Addon]:
|
async def get_all(self) -> list[Addon]:
|
||||||
stmt = select(Addon)
|
stmt = select(Addon)
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from db.models.invoice import Invoice
|
from db.models.invoice import Invoice
|
||||||
|
from db.session import UnitOfWork
|
||||||
from schemas.invoices import InvoiceStatus
|
from schemas.invoices import InvoiceStatus
|
||||||
|
|
||||||
|
|
||||||
class InvoiceRepository:
|
class InvoiceRepository:
|
||||||
def __init__(self, session: AsyncSession) -> None:
|
def __init__(self, uow: UnitOfWork) -> None:
|
||||||
self.session = session
|
self.uow = uow
|
||||||
|
self.session = uow.session
|
||||||
|
|
||||||
async def get_by_id(self, id: int) -> Invoice | None:
|
async def get_by_id(self, id: int) -> Invoice | None:
|
||||||
stmt = select(Invoice).where(Invoice.id == id)
|
stmt = select(Invoice).where(Invoice.id == id)
|
||||||
@@ -21,21 +22,20 @@ class InvoiceRepository:
|
|||||||
|
|
||||||
return list(r.scalars().all())
|
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(
|
obj = Invoice(
|
||||||
creator_id=creator_id,
|
creator_id=creator_id,
|
||||||
|
order_id=order_id,
|
||||||
amount=amount,
|
amount=amount,
|
||||||
status=status,
|
status=status,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.session.add(obj)
|
self.session.add(obj)
|
||||||
await self.session.commit()
|
|
||||||
|
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None:
|
async def update_status_by_id(self, invoice_id: int, status: InvoiceStatus) -> Invoice | None:
|
||||||
invoice = await self.get_by_id(invoice_id)
|
invoice = await self.get_by_id(invoice_id)
|
||||||
invoice.status = status
|
invoice.status = status
|
||||||
await self.session.commit()
|
|
||||||
|
|
||||||
return invoice
|
return invoice
|
||||||
|
|||||||
38
repositories/link_codes.py
Normal file
38
repositories/link_codes.py
Normal 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
75
repositories/orders.py
Normal 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())
|
||||||
@@ -1,12 +1,12 @@
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from db.models.pricing import PricingConfig
|
from db.models.pricing import PricingConfig
|
||||||
|
from db.session import UnitOfWork
|
||||||
|
|
||||||
|
|
||||||
class PricingRepository:
|
class PricingRepository:
|
||||||
def __init__(self, session: AsyncSession) -> None:
|
def __init__(self, uow: UnitOfWork) -> None:
|
||||||
self.session = session
|
self.session = uow.session
|
||||||
|
|
||||||
async def get(self) -> PricingConfig | None:
|
async def get(self) -> PricingConfig | None:
|
||||||
stmt = select(PricingConfig).where(PricingConfig.id == 1)
|
stmt = select(PricingConfig).where(PricingConfig.id == 1)
|
||||||
|
|||||||
61
repositories/service_notifications.py
Normal file
61
repositories/service_notifications.py
Normal 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
|
||||||
16
repositories/service_signatures.py
Normal file
16
repositories/service_signatures.py
Normal 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()
|
||||||
@@ -1,13 +1,14 @@
|
|||||||
from sqlalchemy import func, select
|
from sqlalchemy import func, select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from db.models import Session
|
from db.models import Session
|
||||||
|
from db.session import UnitOfWork
|
||||||
from schemas.providers import ProvidersType
|
from schemas.providers import ProvidersType
|
||||||
|
|
||||||
|
|
||||||
class SessionsRepository:
|
class SessionsRepository:
|
||||||
def __init__(self, session: AsyncSession) -> None:
|
def __init__(self, uow: UnitOfWork) -> None:
|
||||||
self.session = session
|
self.uow = uow
|
||||||
|
self.session = uow.session
|
||||||
|
|
||||||
async def get_session_by_id(self, id: int) -> Session | None:
|
async def get_session_by_id(self, id: int) -> Session | None:
|
||||||
stmt = select(Session).where(Session.id == id)
|
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:
|
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)
|
obj = Session(user_id=user_id, refresh_token_hash=refresh_token_hash, source=iss)
|
||||||
self.session.add(obj)
|
self.session.add(obj)
|
||||||
await self.session.commit()
|
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
async def revoke(self, token_id: int):
|
async def revoke(self, token_id: int):
|
||||||
session = await self.get_session_by_id(token_id)
|
session = await self.get_session_by_id(token_id)
|
||||||
session.is_revoked = True
|
session.is_revoked = True
|
||||||
session.revoked_at = func.now()
|
session.revoked_at = func.now()
|
||||||
await self.session.commit()
|
|
||||||
return session
|
return session
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from db.models import User
|
from db.models import User
|
||||||
from db.models.transactions import BalanceTransaction, BalanceTxType
|
from db.models.transactions import BalanceTransaction, BalanceTxType
|
||||||
|
from db.session import UnitOfWork
|
||||||
|
|
||||||
|
|
||||||
class UserRepository:
|
class UserRepository:
|
||||||
def __init__(self, session: AsyncSession) -> None:
|
def __init__(self, uow: UnitOfWork) -> None:
|
||||||
self.session = session
|
self.uow = uow
|
||||||
|
self.session = uow.session
|
||||||
|
|
||||||
async def get_user_by_id(self, id: int) -> User | None:
|
async def get_user_by_id(self, id: int) -> User | None:
|
||||||
stmt = select(User).where(User.id == id)
|
stmt = select(User).where(User.id == id)
|
||||||
@@ -24,21 +25,26 @@ class UserRepository:
|
|||||||
res = await self.session.execute(stmt)
|
res = await self.session.execute(stmt)
|
||||||
return res.scalar_one_or_none()
|
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(
|
async def create(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
username: str | None = None,
|
username: str | None = None,
|
||||||
hashed_password: str | None = None,
|
hashed_password: str | None = None,
|
||||||
telegram_id: int | None = None,
|
telegram_id: int | None = None,
|
||||||
|
referal_id: int | None = None,
|
||||||
) -> User:
|
) -> User:
|
||||||
obj = User(
|
obj = User(
|
||||||
username=username,
|
username=username,
|
||||||
hashed_password=hashed_password,
|
hashed_password=hashed_password,
|
||||||
telegram_id=telegram_id,
|
telegram_id=telegram_id,
|
||||||
|
referal_id=referal_id,
|
||||||
)
|
)
|
||||||
self.session.add(obj)
|
self.session.add(obj)
|
||||||
await self.session.commit()
|
|
||||||
|
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
async def increase_balance(
|
async def increase_balance(
|
||||||
@@ -60,5 +66,8 @@ class UserRepository:
|
|||||||
self.session.add(obj)
|
self.session.add(obj)
|
||||||
|
|
||||||
user.balance += amount
|
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
|
return user
|
||||||
|
|||||||
@@ -8,3 +8,10 @@ asyncpg>=0.31.0
|
|||||||
alembic>=1.18.0
|
alembic>=1.18.0
|
||||||
aiohttp>=3.14.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
|
||||||
@@ -1,8 +1,21 @@
|
|||||||
from fastapi import APIRouter
|
from fastapi import APIRouter
|
||||||
|
|
||||||
from .auth import router as auth_router
|
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 .orders import router as orders_router
|
||||||
from .payments import payment_routers
|
from .payments import payment_router
|
||||||
from .plans import router as plans_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,
|
||||||
|
]
|
||||||
|
|||||||
@@ -1,12 +1,15 @@
|
|||||||
from fastapi import APIRouter, Depends, HTTPException
|
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 core.secrets import (
|
||||||
from db.session import get_db
|
estimate_password_strength,
|
||||||
|
hash_password,
|
||||||
|
hash_refresh_token,
|
||||||
|
verify_password,
|
||||||
|
)
|
||||||
|
from db.session import UnitOfWork, get_uow
|
||||||
from repositories.sessions import SessionsRepository
|
from repositories.sessions import SessionsRepository
|
||||||
from repositories.users import UserRepository
|
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.providers import ProvidersType
|
||||||
from schemas.registration import UserRegistration
|
from schemas.registration import UserRegistration
|
||||||
from schemas.user import UserInfo
|
from schemas.user import UserInfo
|
||||||
@@ -16,9 +19,10 @@ from services.users import authorize_user
|
|||||||
router = APIRouter(prefix="/auth")
|
router = APIRouter(prefix="/auth")
|
||||||
|
|
||||||
|
|
||||||
@router.post("/signup")
|
@router.post("/signup", response_model=AuthenticatedUser)
|
||||||
async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db)):
|
async def signup(req: UserRegistration, uow: UnitOfWork = Depends(get_uow)) -> AuthenticatedUser:
|
||||||
users_repo = UserRepository(session)
|
users_repo = UserRepository(uow)
|
||||||
|
sessions_repo = SessionsRepository(uow)
|
||||||
|
|
||||||
if req.provider == "credentials":
|
if req.provider == "credentials":
|
||||||
if not req.username or not req.password:
|
if not req.username or not req.password:
|
||||||
@@ -27,19 +31,40 @@ async def signup(req: UserRegistration, session: AsyncSession = Depends(get_db))
|
|||||||
if user:
|
if user:
|
||||||
raise HTTPException(status_code=409, detail="User already exists")
|
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)
|
password_hash = hash_password(req.password)
|
||||||
user = await users_repo.create(username=req.username, hashed_password=password_hash)
|
referal_id = None
|
||||||
return JSONResponse(
|
if req.referal_code:
|
||||||
UserInfo(username=user.username, telegram_id=user.telegram_id).model_dump(),
|
referal = await users_repo.get_user_by_ref_code(req.referal_code)
|
||||||
status_code=201,
|
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")
|
raise HTTPException(status_code=400, detail="Unsupported provider")
|
||||||
|
|
||||||
|
|
||||||
@router.post("/login", response_model=UserLogin)
|
@router.post("/login", response_model=AuthenticatedUser)
|
||||||
async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
|
async def login(req: UserLoginData, uow: UnitOfWork = Depends(get_uow)):
|
||||||
users_repo = UserRepository(session)
|
users_repo = UserRepository(uow)
|
||||||
sessions_repo = SessionsRepository(session)
|
sessions_repo = SessionsRepository(uow)
|
||||||
if req.provider == "credentials":
|
if req.provider == "credentials":
|
||||||
if not req.username or not req.password:
|
if not req.username or not req.password:
|
||||||
raise HTTPException(status_code=400, detail="Username or password is not provided.")
|
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")
|
raise HTTPException(status_code=401, detail="Invalid password")
|
||||||
|
|
||||||
data = await authorize_user(sessions_repo, user, req.provider)
|
data = await authorize_user(sessions_repo, user, req.provider)
|
||||||
|
await uow.commit()
|
||||||
return data
|
return data
|
||||||
|
|
||||||
if req.provider == "telegram":
|
if req.provider == "telegram":
|
||||||
@@ -62,8 +88,8 @@ async def login(req: UserLoginData, session: AsyncSession = Depends(get_db)):
|
|||||||
|
|
||||||
|
|
||||||
@router.post("/refresh", response_model=UserTokens)
|
@router.post("/refresh", response_model=UserTokens)
|
||||||
async def refresh(refresh_token: str, iss: ProvidersType, session: AsyncSession = Depends(get_db)):
|
async def refresh(refresh_token: str, iss: ProvidersType, uow: UnitOfWork = Depends(get_uow)):
|
||||||
sessions_repo = SessionsRepository(session)
|
sessions_repo = SessionsRepository(uow)
|
||||||
|
|
||||||
token_hash = hash_refresh_token(refresh_token)
|
token_hash = hash_refresh_token(refresh_token)
|
||||||
token_entry = await sessions_repo.get_session_by_hash(token_hash)
|
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.")
|
raise HTTPException(status_code=401, detail="Refresh token is invalid.")
|
||||||
|
|
||||||
key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss)
|
key_pair = await refresh_token_rotation(sessions_repo, token_entry, iss)
|
||||||
|
await uow.commit()
|
||||||
return UserTokens(
|
return UserTokens(
|
||||||
access_token=key_pair.access_token,
|
access_token=key_pair.access_token,
|
||||||
refresh_token=key_pair.refresh_token,
|
refresh_token=key_pair.refresh_token,
|
||||||
|
|||||||
8
routes/health.py
Normal file
8
routes/health.py
Normal file
@@ -0,0 +1,8 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/health")
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/")
|
||||||
|
async def healthcheck():
|
||||||
|
return "OK"
|
||||||
6
routes/internal/__init__.py
Normal file
6
routes/internal/__init__.py
Normal 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)
|
||||||
60
routes/internal/renewal.py
Normal file
60
routes/internal/renewal.py
Normal 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
63
routes/link_codes.py
Normal 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,
|
||||||
|
)
|
||||||
@@ -1,41 +1,109 @@
|
|||||||
|
import math
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from config import cfg
|
from config import cfg
|
||||||
from core.deps import get_auth_context, get_pally_client
|
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 external.pally import PallyClient
|
||||||
from repositories import AddonsRepository, PricingRepository
|
from repositories import AddonsRepository, PricingRepository
|
||||||
from repositories.invoices import InvoiceRepository
|
from repositories.invoices import InvoiceRepository
|
||||||
|
from repositories.orders import OrderRepository
|
||||||
|
from schemas.checkout import CheckoutResponse
|
||||||
from schemas.dto import AuthContext
|
from schemas.dto import AuthContext
|
||||||
from schemas.invoices import InvoiceResponse, InvoiceStatus
|
from schemas.invoices import InvoiceStatus
|
||||||
from schemas.plans import OrderDetails
|
from schemas.plans import OrderDetails
|
||||||
from services.plans import calculate_price, get_pricing_model
|
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 = APIRouter(prefix="/orders")
|
||||||
|
|
||||||
|
|
||||||
@router.post("/checkout", response_model=InvoiceResponse, status_code=201)
|
@router.post("/checkout", response_model=CheckoutResponse, status_code=201)
|
||||||
async def checkout(
|
async def checkout(
|
||||||
order: OrderDetails,
|
order: OrderDetails,
|
||||||
ctx: AuthContext = Depends(get_auth_context),
|
ctx: AuthContext = Depends(get_auth_context),
|
||||||
session: AsyncSession = Depends(get_db),
|
uow: UnitOfWork = Depends(get_uow),
|
||||||
pally: PallyClient = Depends(get_pally_client),
|
pally: PallyClient = Depends(get_pally_client),
|
||||||
):
|
):
|
||||||
addons_repo = AddonsRepository(session)
|
addons_repo = AddonsRepository(uow)
|
||||||
pricing_repo = PricingRepository(session)
|
pricing_repo = PricingRepository(uow)
|
||||||
invoices_repo = InvoiceRepository(session)
|
invoices_repo = InvoiceRepository(uow)
|
||||||
|
orders_repo = OrderRepository(uow)
|
||||||
|
|
||||||
pricing = await get_pricing_model(addons_repo, pricing_repo)
|
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(
|
order_entry = await orders_repo.create(
|
||||||
creator_id=ctx.user.id, amount=price, status=InvoiceStatus.ACTIVE
|
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):
|
bill = await pally.bills.create(
|
||||||
raise HTTPException(500, detail="Failed to create an invoice.")
|
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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,3 +1,6 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
|
||||||
from .pally import router as pally_router
|
from .pally import router as pally_router
|
||||||
|
|
||||||
payment_routers = [pally_router]
|
payment_router = APIRouter(prefix="/payments")
|
||||||
|
payment_router.include_router(pally_router)
|
||||||
|
|||||||
@@ -1,28 +1,24 @@
|
|||||||
# ruff: noqa: N803
|
# ruff: noqa: N803
|
||||||
import hashlib
|
|
||||||
import hmac
|
|
||||||
import logging
|
import logging
|
||||||
import math
|
|
||||||
|
|
||||||
from fastapi import Depends, Form, HTTPException
|
from fastapi import Depends, Form, HTTPException
|
||||||
from fastapi.routing import APIRouter
|
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.models.transactions import BalanceTxType
|
||||||
|
from db.session import UnitOfWork, get_uow
|
||||||
from external.pally import BillStatus
|
from external.pally import BillStatus
|
||||||
from repositories.invoices import InvoiceRepository
|
from repositories.invoices import InvoiceRepository
|
||||||
from repositories.users import UserRepository
|
from repositories.users import UserRepository
|
||||||
from schemas.invoices import InvoiceStatus
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/result")
|
@router.post("/result")
|
||||||
async def pally_callback( # noqa: PLR0911
|
async def pally_callback(
|
||||||
*,
|
*,
|
||||||
InvId: str = Form(...),
|
InvId: str = Form(...),
|
||||||
OutSum: str = Form(...),
|
OutSum: str = Form(...),
|
||||||
@@ -32,7 +28,6 @@ async def pally_callback( # noqa: PLR0911
|
|||||||
CurrencyIn: str = Form(...),
|
CurrencyIn: str = Form(...),
|
||||||
custom: str | None = Form(None),
|
custom: str | None = Form(None),
|
||||||
SignatureValue: str = Form(...),
|
SignatureValue: str = Form(...),
|
||||||
# Optional fields for additional information
|
|
||||||
AccountType: str | None = Form(None),
|
AccountType: str | None = Form(None),
|
||||||
AccountNumber: str | None = Form(None),
|
AccountNumber: str | None = Form(None),
|
||||||
BalanceAmount: str | None = Form(None),
|
BalanceAmount: str | None = Form(None),
|
||||||
@@ -43,12 +38,8 @@ async def pally_callback( # noqa: PLR0911
|
|||||||
PayerComment: str | None = Form(None),
|
PayerComment: str | None = Form(None),
|
||||||
ErrorCode: int | None = Form(None),
|
ErrorCode: int | None = Form(None),
|
||||||
ErrorMessage: str | 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(
|
logger.info(
|
||||||
"Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, "
|
"Pally webhook received - InvId: %s, OutSum: %s, Commission: %s, TrsId: %s, Status: %s, "
|
||||||
"CurrencyIn: %s, custom: %s, BalanceAmount: %s, SignatureValue: %s",
|
"CurrencyIn: %s, custom: %s, BalanceAmount: %s, SignatureValue: %s",
|
||||||
@@ -72,29 +63,18 @@ async def pally_callback( # noqa: PLR0911
|
|||||||
TrsId,
|
TrsId,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Validate signature
|
if not validate_pally_signature(OutSum, InvId, SignatureValue):
|
||||||
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):
|
|
||||||
logger.critical(
|
logger.critical(
|
||||||
"SECURITY ALERT: Invalid signature for TrsId %s - Expected: %s, Received: %s",
|
"SECURITY ALERT: Invalid signature for TrsId %s",
|
||||||
TrsId,
|
TrsId,
|
||||||
expected_signature,
|
|
||||||
SignatureValue,
|
|
||||||
)
|
)
|
||||||
raise HTTPException(403, detail="Invalid signature.")
|
raise HTTPException(403, detail="Invalid signature.")
|
||||||
|
|
||||||
# Only process successful payments
|
|
||||||
if Status != BillStatus.SUCCESS:
|
if Status != BillStatus.SUCCESS:
|
||||||
logger.info("Bill %s skipped: status=%s", TrsId, Status)
|
logger.info("Bill %s skipped: status=%s", TrsId, Status)
|
||||||
return "OK"
|
return "OK"
|
||||||
|
|
||||||
logger.info("Processing successfully paid bill %s", TrsId)
|
invoice_id_str = InvId
|
||||||
|
|
||||||
# Validate bill ID (from InvId field, which contains the order_id from bill creation)
|
|
||||||
if not invoice_id_str or not invoice_id_str.isdigit():
|
if not invoice_id_str or not invoice_id_str.isdigit():
|
||||||
logger.critical(
|
logger.critical(
|
||||||
"Invalid or non-numeric bill ID in InvId field for TrsId %s: '%s'",
|
"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"
|
return "OK"
|
||||||
|
|
||||||
# Find bill in database
|
amount = int(float(BalanceAmount)) if BalanceAmount is not None else int(float(OutSum))
|
||||||
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"
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Handle fee scenarios: use BalanceAmount if available (net amount after fees),
|
await process_subscription_purchase(
|
||||||
# otherwise use OutSum (gross amount paid by customer)
|
uow,
|
||||||
if BalanceAmount is not None:
|
invoice_id=int(invoice_id_str),
|
||||||
# Customer pays fees - BalanceAmount is the net amount credited to merchant
|
trs_id=TrsId,
|
||||||
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,
|
|
||||||
amount=amount,
|
amount=amount,
|
||||||
tx_type=BalanceTxType.DEPOSIT,
|
|
||||||
description=f"payment via PALLY (TrsId: {TrsId})",
|
|
||||||
)
|
)
|
||||||
|
except Exception:
|
||||||
# Process referral bonus
|
# The payment is confirmed, so never leave the user without the money
|
||||||
user = invoice.creator
|
# if order/subscription provisioning fails.
|
||||||
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:
|
|
||||||
logger.exception(
|
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,
|
TrsId,
|
||||||
invoice.id,
|
|
||||||
invoice.creator_id,
|
|
||||||
str(e),
|
|
||||||
)
|
)
|
||||||
# Don't return early - still mark as success to prevent retries
|
await uow.rollback()
|
||||||
# The balance operation might have partially succeeded
|
|
||||||
|
|
||||||
# Update bill status to success
|
invoice_repo = InvoiceRepository(uow)
|
||||||
await invoice_repo.update_status_by_id(int(invoice.id), status=InvoiceStatus.PAID)
|
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"
|
return "OK"
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from fastapi import APIRouter, Depends
|
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.addons import AddonsRepository
|
||||||
from repositories.pricing import PricingRepository
|
from repositories.pricing import PricingRepository
|
||||||
from schemas.plans import PricingPlans
|
from schemas.plans import PricingPlans
|
||||||
@@ -11,9 +10,9 @@ router = APIRouter(prefix="/plans")
|
|||||||
|
|
||||||
|
|
||||||
@router.get("/", response_model=PricingPlans)
|
@router.get("/", response_model=PricingPlans)
|
||||||
async def get_plans(session: AsyncSession = Depends(get_db)):
|
async def get_plans(uow: UnitOfWork = Depends(get_uow)):
|
||||||
addons_repo = AddonsRepository(session)
|
addons_repo = AddonsRepository(uow)
|
||||||
pricing_repo = PricingRepository(session)
|
pricing_repo = PricingRepository(uow)
|
||||||
|
|
||||||
res = await get_pricing_model(addons_repo, pricing_repo)
|
res = await get_pricing_model(addons_repo, pricing_repo)
|
||||||
return res
|
return res
|
||||||
|
|||||||
97
routes/users.py
Normal file
97
routes/users.py
Normal 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
15
schemas/checkout.py
Normal 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
5
schemas/common.py
Normal file
@@ -0,0 +1,5 @@
|
|||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class OperationData(BaseModel):
|
||||||
|
success: bool = Field(False)
|
||||||
8
schemas/devices.py
Normal file
8
schemas/devices.py
Normal 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()
|
||||||
@@ -14,5 +14,14 @@ class KeyPair:
|
|||||||
@dataclass
|
@dataclass
|
||||||
class AuthContext:
|
class AuthContext:
|
||||||
user: User
|
user: User
|
||||||
auth_method: Literal["jwt"]
|
auth_method: Literal["jwt", "service"]
|
||||||
service: str | None = None
|
service: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ServiceIdentity:
|
||||||
|
"""
|
||||||
|
Proves that request is coming from verified source, no user context provided.
|
||||||
|
"""
|
||||||
|
|
||||||
|
service: str
|
||||||
|
|||||||
31
schemas/enums.py
Normal file
31
schemas/enums.py
Normal 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"
|
||||||
@@ -3,8 +3,17 @@ from pydantic import BaseModel, Field
|
|||||||
from core.exp import get_exp
|
from core.exp import get_exp
|
||||||
from schemas.providers import ProvidersType
|
from schemas.providers import ProvidersType
|
||||||
|
|
||||||
|
type NumericDate = int | float
|
||||||
|
|
||||||
class JWTPayload(BaseModel):
|
|
||||||
|
class UserJWTPayload(BaseModel):
|
||||||
sub: str = Field(description="User ID")
|
sub: str = Field(description="User ID")
|
||||||
iss: ProvidersType = Field(description="Issuer")
|
iss: ProvidersType = Field(description="Issuer")
|
||||||
exp: float = Field(default_factory=get_exp)
|
exp: NumericDate = Field(default_factory=get_exp)
|
||||||
|
|
||||||
|
|
||||||
|
class ServiceJWTPayload(BaseModel):
|
||||||
|
iss: str = Field(description="Service name")
|
||||||
|
iat: NumericDate = Field(description="Issued at")
|
||||||
|
exp: NumericDate = Field(description="Expiration date")
|
||||||
|
acting_as: str = Field(description="Subject of issued call, formatted as `service:id`")
|
||||||
|
|||||||
20
schemas/link_codes.py
Normal file
20
schemas/link_codes.py
Normal 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()
|
||||||
@@ -16,7 +16,7 @@ class UserLoginData(BaseModel):
|
|||||||
telegram: TelegramData | None = None
|
telegram: TelegramData | None = None
|
||||||
|
|
||||||
|
|
||||||
class UserLogin(BaseModel):
|
class AuthenticatedUser(BaseModel):
|
||||||
access_token: str
|
access_token: str
|
||||||
refresh_token: str
|
refresh_token: str
|
||||||
expires_at: float
|
expires_at: float
|
||||||
|
|||||||
24
schemas/notifications.py
Normal file
24
schemas/notifications.py
Normal 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()
|
||||||
@@ -6,6 +6,7 @@ class AddonData(BaseModel):
|
|||||||
name: str
|
name: str
|
||||||
price: float
|
price: float
|
||||||
free_threshold: int
|
free_threshold: int
|
||||||
|
is_enabled: bool
|
||||||
|
|
||||||
|
|
||||||
class PricingPlans(BaseModel):
|
class PricingPlans(BaseModel):
|
||||||
|
|||||||
@@ -9,4 +9,6 @@ class UserRegistration(BaseModel):
|
|||||||
username: str | None = Field(None)
|
username: str | None = Field(None)
|
||||||
password: str | None = Field(None)
|
password: str | None = Field(None)
|
||||||
|
|
||||||
|
referal_code: str | None = Field(None)
|
||||||
|
|
||||||
provider: ProvidersType
|
provider: ProvidersType
|
||||||
|
|||||||
@@ -1,6 +1,19 @@
|
|||||||
|
from datetime import datetime
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
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):
|
class UserInfo(BaseModel):
|
||||||
username: str | None = Field()
|
username: str | None = Field(None)
|
||||||
telegram_id: str | None = Field()
|
telegram_id: int | None = Field(None)
|
||||||
|
referal_code: str = Field()
|
||||||
|
bonus_balance: float = Field()
|
||||||
|
|||||||
59
services/notifications.py
Normal file
59
services/notifications.py
Normal 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
171
services/payments.py
Normal 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)
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
from repositories.addons import AddonsRepository
|
from repositories.addons import AddonsRepository
|
||||||
from repositories.pricing import PricingRepository
|
from repositories.pricing import PricingRepository
|
||||||
from schemas.plans import AddonData, OrderDetails, PricingPlans
|
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):
|
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 = await addons_repo.get_all()
|
||||||
addons_data = [
|
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
|
for a in addons
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -21,4 +28,9 @@ async def calculate_price(
|
|||||||
*, addons_repo: AddonsRepository, order: OrderDetails, pricing: PricingPlans
|
*, addons_repo: AddonsRepository, order: OrderDetails, pricing: PricingPlans
|
||||||
) -> float:
|
) -> float:
|
||||||
addons = [await addons_repo.get_by_id(a) for a in order.addons]
|
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
126
services/rw_sync.py
Normal 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
197
services/subscriptions.py
Normal 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)
|
||||||
@@ -1,23 +1,28 @@
|
|||||||
from core.secrets import generate_pair, hash_refresh_token
|
from core.secrets import generate_pair, hash_refresh_token
|
||||||
from db.models.users import User
|
from db.models.users import User
|
||||||
from repositories.sessions import SessionsRepository
|
from repositories.sessions import SessionsRepository
|
||||||
from schemas.login import UserLogin
|
from schemas.login import AuthenticatedUser
|
||||||
from schemas.providers import ProvidersType
|
from schemas.providers import ProvidersType
|
||||||
from schemas.user import UserInfo
|
from schemas.user import UserInfo
|
||||||
|
|
||||||
|
|
||||||
async def authorize_user(
|
async def authorize_user(
|
||||||
sessions_repo: SessionsRepository, user: User, iss: ProvidersType
|
sessions_repo: SessionsRepository, user: User, iss: ProvidersType
|
||||||
) -> UserLogin:
|
) -> AuthenticatedUser:
|
||||||
key_pair = generate_pair(user.id, iss)
|
key_pair = generate_pair(user.id, iss)
|
||||||
|
|
||||||
refresh_token_hash = hash_refresh_token(key_pair.refresh_token)
|
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)
|
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,
|
access_token=key_pair.access_token,
|
||||||
refresh_token=key_pair.refresh_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,
|
expires_at=key_pair.expires_at,
|
||||||
)
|
)
|
||||||
|
|||||||
17
tests/conftest.py
Normal file
17
tests/conftest.py
Normal 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
127
tests/test_auth_routes.py
Normal 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."
|
||||||
137
tests/test_link_code_routes.py
Normal file
137
tests/test_link_code_routes.py
Normal 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)
|
||||||
47
tests/test_notifications.py
Normal file
47
tests/test_notifications.py
Normal 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,
|
||||||
|
]
|
||||||
106
tests/test_services_and_payments.py
Normal file
106
tests/test_services_and_payments.py
Normal 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
|
||||||
131
tests/test_user_and_plan_routes.py
Normal file
131
tests/test_user_and_plan_routes.py
Normal 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
|
||||||
Reference in New Issue
Block a user