Files
finance-dashboard/tests/test_trading.py
T

210 lines
12 KiB
Python

import csv
from datetime import date
from decimal import Decimal
import io
import os
from pathlib import Path
import secrets
import sqlite3
import sys
import tempfile
import unittest
from unittest.mock import patch
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / 'app'))
from fastapi.testclient import TestClient
from main import app
from database import connect, initialize
from services import trading_service as service
from services.asset_service import create_asset
from trading_models import TradingValidationError
class TradingTests(unittest.TestCase):
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.path = Path(self.tmp.name) / 'finance.db'
self.token = secrets.token_urlsafe(32)
self.env = patch.dict(os.environ, {'FINANCE_DB_PATH':str(self.path), 'FINANCE_API_TOKEN':self.token})
self.env.start()
self.client = TestClient(app)
self.client.__enter__()
self.asset = create_asset('SpaceX', 'stock')['id']
self.headers = {'Authorization':'Bearer '+self.token}
def tearDown(self):
self.client.__exit__(None,None,None)
self.env.stop()
self.tmp.cleanup()
def trade(self, **changes):
values = dict(date='2026-09-01',asset_id=self.asset,transaction_type='buy',quantity='10',price_per_unit='10',currency='EUR',fees='0',source='manual',strategy_tag=None)
values.update(changes)
return service.save_transaction(values)
def api(self, method, path, **kwargs):
return self.client.request(method,'/api/v1'+path,headers=self.headers,**kwargs)
def test_buy_multiple_average_partial_sale_and_realized(self):
self.trade(fees='2')
self.trade(price_per_unit='20',fees='2',source='savings_plan',strategy_tag='core')
self.trade(transaction_type='sell',quantity='5',price_per_unit='30',fees='1')
position = service.positions()[0]
self.assertEqual(Decimal(position['quantity']),15)
self.assertEqual(position['invested_capital'],'228.00')
self.assertEqual(Decimal(position['average_cost']),Decimal('15.2'))
self.assertEqual(position['realized_profit_loss'],'73.00')
self.assertEqual(position['total_buys'],'304.00')
self.assertEqual(position['total_sells'],'149.00')
self.assertEqual((position['buy_count'],position['sell_count']),(2,1))
stats = service.trading_stats(today=date(2026,9,9))
self.assertEqual(stats['realized_profit_loss_current_year'],'73.00')
self.assertEqual(stats['transactions_total'],3)
source = {r['key']:r['amount'] for r in stats['by_source']}
self.assertEqual(source['manual'],'76.50')
self.assertEqual(source['savings_plan'],'151.50')
for field in ['by_source','by_strategy']:
self.assertEqual(sum(Decimal(r['amount']) for r in stats[field]),Decimal('228'))
def test_roundup_cashback_savings_plan_and_precise_balance(self):
self.trade(quantity='2.600000',price_per_unit='100',strategy_tag='conviction')
roundup = self.trade(quantity='0,092262',price_per_unit='127,24',source='roundup',strategy_tag='conviction')
self.assertEqual(roundup['gross_amount'],'11.74')
detail = service.asset_detail(self.asset)
self.assertEqual(Decimal(detail['position']['quantity']),Decimal('2.692262'))
self.assertEqual(Decimal(detail['entries'][0]['quantity_before']),Decimal('2.600000'))
self.assertEqual(Decimal(detail['entries'][0]['quantity_after']),Decimal('2.692262'))
self.trade(quantity='0.000000000001',price_per_unit='1',source='cashback')
self.trade(quantity='0.1',price_per_unit='100',source='savings_plan',strategy_tag='core')
self.assertEqual(len(service.list_transactions(source='roundup')),1)
self.assertEqual(len(service.list_transactions(source='cashback')),1)
self.assertEqual(len(service.list_transactions(source='savings_plan')),1)
self.assertEqual(len(service.list_transactions(strategy_tag='conviction')),2)
self.assertEqual(Decimal(service.positions()[0]['quantity']),Decimal('2.792262000001'))
def test_oversell_and_backdated_sale(self):
self.trade()
for values in [dict(transaction_type='sell',quantity='10.000000000001'),dict(transaction_type='sell',date='2026-08-01',quantity='1')]:
with self.assertRaises(TradingValidationError):
self.trade(**values)
self.assertEqual(len(service.list_transactions()),1)
def test_edit_delete_validate_full_history_and_rollback(self):
buy = self.trade()
sell = self.trade(transaction_type='sell',quantity='8',date='2026-09-02')
for changes in [{'quantity':'7'}, {'date':'2026-09-03'}, {'transaction_type':'sell'}]:
with self.assertRaises(TradingValidationError):
service.save_transaction(changes,buy['id'])
with self.assertRaises(TradingValidationError):
service.delete_transaction(buy['id'])
self.assertEqual(service.get_transaction(buy['id'])['quantity'],'10')
updated = service.save_transaction({'price_per_unit':'20'},buy['id'])
self.assertEqual(updated['source'],'manual')
self.assertEqual(service.positions()[0]['realized_profit_loss'],'-80.00')
service.delete_transaction(sell['id'])
self.assertEqual(Decimal(service.positions()[0]['quantity']),10)
service.delete_transaction(buy['id'])
self.assertEqual(service.positions(),[])
def test_asset_move_cannot_leave_old_asset_short(self):
second = create_asset('ETF Test','etf')['id']
buy = self.trade()
self.trade(transaction_type='sell',quantity='1')
with self.assertRaises(TradingValidationError):
service.save_transaction({'asset_id':second},buy['id'])
self.assertEqual(service.get_transaction(buy['id'])['asset_id'],self.asset)
def test_full_disposal_resets_cost_and_reopening(self):
self.trade(quantity='3',price_per_unit='0.01',fees='0.01')
self.trade(transaction_type='sell',quantity='1',price_per_unit='0.02')
self.trade(transaction_type='sell',quantity='2',price_per_unit='0.02')
self.assertEqual(service.positions(),[])
position = service.positions(include_closed=True)[0]
self.assertEqual(position['invested_capital'],'0.00')
self.assertEqual(position['realized_profit_loss'],'0.02')
self.trade(quantity='1',price_per_unit='10')
self.assertEqual(Decimal(service.positions()[0]['average_cost']),10)
def test_multi_currency_separation(self):
self.trade()
second = create_asset('USD ETF','etf')['id']
self.trade(asset_id=second,currency='USD',price_per_unit='99')
self.assertEqual(service.trading_stats('EUR')['invested_capital'],'100.00')
self.assertEqual(service.trading_stats('USD')['invested_capital'],'990.00')
with self.assertRaises(TradingValidationError):
self.trade(currency='USD')
self.assertEqual(len(service.positions()),2)
def test_invalid_values(self):
for change in [{'quantity':'0'},{'quantity':'-1'},{'quantity':0.1},{'price_per_unit':'-1'},{'fees':'-1'},{'fees':'0.001'},
{'quantity':'NaN'},{'price_per_unit':'Infinity'},{'quantity':'0.0000000000001'}, {'asset_id':99999},
{'date':'2026-02-30'},{'currency':'EURO'},{'source':'invalid'},{'strategy_tag':'invalid'}]:
with self.subTest(change=change), self.assertRaises(TradingValidationError):
self.trade(**change)
self.assertEqual(service.list_transactions(),[])
def test_api_auth_crud_and_filters(self):
paths = [('GET','/transactions'),('POST','/transactions'),('GET','/transactions/1'),('PATCH','/transactions/1'),('DELETE','/transactions/1'),
('GET','/positions'),('GET','/trading/stats'),('GET','/trading/by-source'),('GET','/trading/by-strategy')]
for method,path in paths:
self.assertEqual(self.client.request(method,'/api/v1'+path).status_code,401)
response = self.api('POST','/transactions',json={'date':'2026-09-01','asset_id':self.asset,'quantity':'0.092262','price_per_unit':'127.24','source':'roundup','strategy_tag':'conviction'})
self.assertEqual(response.status_code,201,response.text)
entry = response.json()
self.assertEqual(entry['quantity'],'0.092262')
self.assertEqual(entry['total_cost'],'11.74')
self.assertEqual(self.api('GET',f"/transactions/{entry['id']}").json(),entry)
response = self.api('PATCH',f"/transactions/{entry['id']}",json={'note':'Test','strategy_tag':None})
self.assertEqual(response.status_code,200,response.text)
self.assertIsNone(response.json()['strategy_tag'])
self.assertEqual(len(self.api('GET','/transactions?source=roundup&strategy_tag=untagged&year=2026&month=9').json()),1)
self.assertEqual(len(self.api('GET','/positions').json()),1)
self.assertEqual(self.api('GET','/trading/stats').json()['invested_capital'],'11.74')
self.assertEqual(self.api('GET','/trading/by-source').json()['currency'],'EUR')
self.assertEqual(self.api('GET','/trading/by-strategy').status_code,200)
self.assertEqual(self.api('POST','/transactions',json={'date':'2026-09-02','asset_id':self.asset,'transaction_type':'sell','quantity':'1','price_per_unit':'1'}).status_code,422)
self.assertEqual(self.api('DELETE',f"/transactions/{entry['id']}").content,b'')
self.assertEqual(self.api('GET',f"/transactions/{entry['id']}").status_code,404)
self.assertEqual(self.api('PATCH','/transactions/999',json={'note':'x'}).status_code,404)
for query in ['source=invalid','month=13','limit=0','strategy_tag=invalid','offset=-1']:
self.assertEqual(self.api('GET','/transactions?'+query).status_code,422)
def test_web_csv_and_income_untouched(self):
with connect() as db:
db.execute("INSERT INTO income_entries(date,asset_id,category,amount) VALUES ('2026-09-01',?,'dividend',4)",(self.asset,))
before = [tuple(r) for r in db.execute('SELECT * FROM income_entries')]
form = dict(date='2026-09-01',asset_id=self.asset,transaction_type='buy',quantity='0,092262',price_per_unit='127,24',currency='EUR',fees='0',source='roundup',strategy_tag='conviction',note='<script>test</script>')
response = self.client.post('/trading/transactions/new',data=form,follow_redirects=False)
self.assertEqual(response.status_code,303,response.text)
for path in ['/','/health','/income','/trading','/trading/positions','/trading/transactions','/trading/transactions/new',f'/trading/assets/{self.asset}','/trading/assets/new']:
response = self.client.get(path)
self.assertEqual(response.status_code,200,(path,response.text))
self.assertIn('&lt;script&gt;',self.client.get('/trading/transactions').text)
self.assertNotIn('<script>test</script>',self.client.get('/trading/transactions').text)
export = self.client.get('/export/trading.csv')
self.assertEqual(export.status_code,200)
self.assertTrue(export.content.startswith(b'\xef\xbb\xbf'))
rows = list(csv.reader(io.StringIO(export.content.decode('utf-8-sig')),delimiter=';'))
self.assertEqual(rows[1][3],'0,092262')
self.assertEqual(rows[1][7],'11,74')
trade = service.list_transactions()[0]
self.assertEqual(self.client.get(f"/trading/transactions/{trade['id']}/edit").status_code,200)
self.assertEqual(self.client.post('/trading/transactions/new',data={**form,'quantity':'0'}).status_code,422)
initialize()
with connect() as db:
self.assertEqual(before,[tuple(r) for r in db.execute('SELECT * FROM income_entries')])
def test_empty_trading_and_schema_migration(self):
self.assertEqual(service.trading_stats()['invested_capital'],'0.00')
self.assertEqual(self.client.get('/trading').status_code,200)
with connect() as db:
before = [tuple(r) for r in db.execute('SELECT * FROM income_entries')]
db.execute('PRAGMA user_version=2')
initialize()
initialize()
with connect() as db:
self.assertEqual(db.execute('PRAGMA user_version').fetchone()[0],3)
self.assertEqual(db.execute('SELECT COUNT(*) FROM transactions').fetchone()[0],0)
self.assertEqual(before,[tuple(r) for r in db.execute('SELECT * FROM income_entries')])
self.assertEqual(db.execute("SELECT COUNT(*) FROM data_migrations WHERE name='2026-09-09-trading-schema'").fetchone()[0],1)