Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
"""Add unique constraint on vault(user_id, symbol)

Revision ID: b2c3d4e5f6a7
Revises: outbox_event_table_rev
Create Date: 2026-08-20

"""
from alembic import op
import sqlalchemy as sa


# revision identifiers, used by Alembic.
revision = "b2c3d4e5f6a7"
down_revision = "outbox_event_table_rev"
branch_labels = None
depends_on = None


def upgrade() -> None:
"""Deduplicates vault rows then adds a unique constraint on (user_id, symbol)."""
conn = op.get_bind()

# Deduplicate: for each (user_id, symbol) pair keep only the row with the
# latest updated_at and sum all amounts into it, then delete the rest.
conn.execute(sa.text("""
WITH ranked AS (
SELECT id,
user_id,
symbol,
amount,
updated_at,
ROW_NUMBER() OVER (
PARTITION BY user_id, symbol ORDER BY updated_at DESC, id
) AS rn
FROM vault
),
totals AS (
SELECT user_id,
symbol,
SUM(amount::numeric) AS total_amount
FROM vault
GROUP BY user_id, symbol
)
UPDATE vault v
SET amount = t.total_amount::text
FROM totals t
WHERE v.user_id = t.user_id
AND v.symbol = t.symbol
AND v.id IN (
SELECT id FROM ranked WHERE rn > 1
);

DELETE FROM vault
WHERE id IN (
SELECT id FROM (
SELECT id,
ROW_NUMBER() OVER (
PARTITION BY user_id, symbol ORDER BY updated_at DESC, id
) AS rn
FROM vault
) sub
WHERE rn > 1
);
"""))

op.create_unique_constraint(
"uq_vault_user_symbol", "vault", ["user_id", "symbol"]
)


def downgrade() -> None:
"""Drops the unique constraint on vault(user_id, symbol)."""
op.drop_constraint("uq_vault_user_symbol", "vault", type_="unique")
58 changes: 44 additions & 14 deletions quantara/web_app/db/crud/deposit.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,13 @@
"""

import logging
import uuid
from decimal import Decimal
from typing import TypeVar

from sqlalchemy import Numeric, cast, func
from sqlalchemy.dialects.postgresql import insert as pg_insert

from web_app.db.models import Base, User, Vault

from .base import DBConnector
Expand All @@ -20,19 +24,50 @@ class DepositDBConnector(DBConnector):
Provides database connection and operations management for the Vault model.
"""

def upsert_vault(self, user_id: uuid.UUID, symbol: str, amount: str) -> Vault:
"""
Atomically inserts a new vault row or adds to the existing balance.

Uses PostgreSQL INSERT ... ON CONFLICT DO UPDATE so that concurrent
calls for the same (user_id, symbol) never lose updates.

:param user_id: UUID of the user
:param symbol: Token symbol or address
:param amount: Amount to add (as string)

:return: Vault instance (existing row updated or newly inserted)
"""
with self.Session() as db:
stmt = (
pg_insert(Vault)
.values(user_id=user_id, symbol=symbol, amount=amount)
.on_conflict_do_update(
constraint="uq_vault_user_symbol",
set_={
"amount": cast(Vault.amount, Numeric)
+ cast(amount, Numeric),
"updated_at": func.now(),
},
)
.returning(Vault)
)
result = db.execute(stmt)
vault = result.scalar_one()
db.commit()
db.refresh(vault)
return vault

def create_vault(self, user: User, symbol: str, amount: str) -> Vault:
"""
Creates a new vault instance
Creates a new vault instance or updates existing balance atomically.

:param user: A user model instance
:param symbol: Token symbol or address
:param amount: An amount in string

:return: Vault
"""
vault = Vault(user_id=user.id, symbol=symbol, amount=amount)
self.write_to_db(vault)
return vault
return self.upsert_vault(user.id, symbol, amount)

def get_vault(self, wallet_id: str, symbol: str) -> Vault | None:
"""
Expand All @@ -53,23 +88,18 @@ def get_vault(self, wallet_id: str, symbol: str) -> Vault | None:

def add_vault_balance(self, wallet_id: str, symbol: str, amount: str) -> Vault:
"""
Adds balance to user vault for token symbol
Adds balance to user vault for token symbol atomically.

:param wallet_id: Wallet id of user
:param symbol: Token symbol or address
:param amount: An amount in string

:return: Updated Vault instance
"""
vault = self.get_vault(wallet_id, symbol)
if not vault:
raise ValueError("Vault not found")
with self.Session() as db:
new_amount = Decimal(vault.amount) + Decimal(amount)
db.query(Vault).filter_by(id=vault.id).update(amount=str(new_amount))
db.commit()
vault = self.get_vault(wallet_id, symbol)
return vault
user = self.get_object_by_field(User, "wallet_id", wallet_id)
if not user:
raise ValueError("User not found")
return self.upsert_vault(user.id, symbol, amount)

def get_vault_balance(self, wallet_id: str, symbol: str) -> str | None:
"""
Expand Down
4 changes: 4 additions & 0 deletions quantara/web_app/db/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,10 @@ class Vault(Base):
DateTime, nullable=False, default=func.now(), onupdate=func.now()
)

__table_args__ = (
UniqueConstraint("user_id", "symbol", name="uq_vault_user_symbol"),
)


class TransactionStatus(PyEnum):
"""
Expand Down
110 changes: 52 additions & 58 deletions quantara/web_app/tests/db/deposit_db_tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,59 +3,64 @@
"""

from decimal import Decimal
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
import uuid

import pytest
from sqlalchemy import create_engine, StaticPool
from sqlalchemy.orm import sessionmaker

from web_app.db.crud import DepositDBConnector
from web_app.db.models import User, Vault
from web_app.db.models import Base, User, Vault


@pytest.fixture
def deposit_db_connector_fixture(mock_db_connector):
"""
Fixture to provide a mocked DepositDBConnector instance using mock_db_connector.
"""
connector = DepositDBConnector()
connector.Session = mock_db_connector.Session
connector.get_object_by_field = mock_db_connector.get_object_by_field
connector.write_to_db = mock_db_connector.write_to_db
def db_session_factory():
"""Create an in-memory SQLite database for testing."""
engine = create_engine(
"sqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine)
return Session


@pytest.fixture
def deposit_connector(db_session_factory):
"""Provide a DepositDBConnector backed by an in-memory SQLite database."""
connector = object.__new__(DepositDBConnector)
connector.Session = db_session_factory
return connector


@pytest.fixture
def mock_user_fixture():
def mock_user():
"""
Mocked User instance.
"""
return User(id="user123", wallet_id="wallet123")
return User(id=uuid.uuid4(), wallet_id="wallet123")


@pytest.fixture
def mock_vault_fixture():
def mock_vault():
"""
Mocked Vault instance.
"""
return Vault(id="vault123", user_id="user123", symbol="ETH", amount="100.00")
return Vault(id=uuid.uuid4(), user_id=uuid.uuid4(), symbol="ETH", amount="100.00")


class TestCreateVault:
"""
Tests for creating a vault using DepositDBConnector.
"""

def test_create_vault_success(
self,
deposit_db_connector: DepositDBConnector,
mock_user: User,
mock_db_connector,
):
def test_create_vault_success(self, deposit_connector, mock_user):
"""
Test successful creation of a vault using fixtures.
Test successful creation of a vault via upsert.
"""
mock_db_connector.write_to_db = MagicMock() # Use mock_db_connector's method

vault = deposit_db_connector.create_vault(
vault = deposit_connector.create_vault(
user=mock_user,
symbol="BTC",
amount="50.00",
Expand All @@ -64,17 +69,23 @@
assert vault.symbol == "BTC"
assert vault.amount == "50.00"
assert vault.user_id == mock_user.id
mock_db_connector.write_to_db.assert_called_once_with(vault)

def test_create_vault_failure_invalid_user(
self,
deposit_db_connector: DepositDBConnector,
):
def test_create_vault_idempotent(self, deposit_connector, mock_user):
"""
Test that calling create_vault twice adds amounts (upsert behaviour).
"""
deposit_connector.create_vault(user=mock_user, symbol="BTC", amount="50.00")
vault = deposit_connector.create_vault(
user=mock_user, symbol="BTC", amount="25.00"
)
assert Decimal(vault.amount) == Decimal("75.00")

def test_create_vault_failure_invalid_user(self, deposit_connector):
"""
Test failure when creating a vault with an invalid user.
"""
with pytest.raises(ValueError, match="Invalid user provided"):
deposit_db_connector.create_vault(
with pytest.raises((AttributeError, TypeError)):
deposit_connector.create_vault(
user=None,
symbol="BTC",
amount="50.00",
Expand All @@ -86,41 +97,24 @@
Tests for adding to a vault's balance using DepositDBConnector.
"""

def test_add_balance_success(
self,
deposit_db_connector: DepositDBConnector,
mock_vault: Vault,
mock_db_connector,
):
def test_add_balance_success(self, deposit_connector, mock_user):
"""
Test successfully adding to a vault's balance using fixtures.
Test successfully adding to a vault's balance.
"""
mock_db_connector.get_object_by_field = MagicMock(return_value=mock_vault)
mock_db_connector.Session().query().filter_by().update = MagicMock()

deposit_db_connector.add_vault_balance(
wallet_id="wallet123",
deposit_connector.upsert_vault(mock_user.id, "ETH", "100.00")
vault = deposit_connector.add_vault_balance(
wallet_id=mock_user.wallet_id,
symbol="ETH",
amount="50.00",
)
assert Decimal(vault.amount) == Decimal("150.00")

updated_amount = Decimal(mock_vault.amount) + Decimal("50.00")
mock_db_connector.Session().query().filter_by().update.assert_called_once_with(
{"amount": str(updated_amount)}
)

def test_add_balance_failure_vault_not_found(
self,
deposit_db_connector: DepositDBConnector,
mock_db_connector,
):
def test_add_balance_failure_vault_not_found(self, deposit_connector):
"""
Test failure when adding to a vault balance that doesn't exist.
Test failure when adding to a vault balance for a non-existent user.
"""
mock_db_connector.get_object_by_field = MagicMock(return_value=None)

with pytest.raises(ValueError, match="Vault not found"):
deposit_db_connector.add_vault_balance(
with pytest.raises(ValueError, match="User not found"):
deposit_connector.add_vault_balance(
wallet_id="invalid_wallet",
symbol="ETH",
amount="50.00",
Expand Down
Loading
Loading