from flask import Blueprint, render_template, redirect, url_for, flash, request, jsonify from flask_login import login_required from flask_wtf import FlaskForm from wtforms import StringField, SelectField, DecimalField, DateField, TextAreaField, SubmitField from wtforms.validators import DataRequired, Optional, NumberRange, Length from app.extensions import db from app.models.investment import Investment, InvestmentTransaction from app.services.investment_service import ( get_portfolio_summary, update_prices, fetch_price, fetch_price_history, fetch_day_change, get_price_alerts, ASSET_COLORS, ASSET_TYPE_LABELS ) from datetime import date from decimal import Decimal investments_bp = Blueprint('investments', __name__, url_prefix='/investments') ASSET_TYPES = [ ('stock', 'Stock'), ('etf', 'ETF'), ('crypto', 'Crypto'), ('real_estate', 'Real Estate'), ('bond', 'Bond'), ('cash', 'Cash'), ('other', 'Other'), ] TXN_TYPES = [ ('buy', 'Buy'), ('sell', 'Sell'), ('dividend', 'Dividend'), ('split', 'Split'), ] class InvestmentForm(FlaskForm): asset_name = StringField('Asset Name', validators=[DataRequired(), Length(1, 100)]) ticker = StringField('Ticker Symbol', validators=[Optional(), Length(max=20)], description='e.g. AAPL, BTC-USD, VNM') asset_type = SelectField('Asset Type', choices=ASSET_TYPES, validators=[DataRequired()]) shares = DecimalField('Shares / Units', validators=[Optional(), NumberRange(min=0)], places=8, default=Decimal('0')) avg_cost_basis = DecimalField('Avg Cost per Share', validators=[Optional(), NumberRange(min=0)], places=4, default=Decimal('0')) notes = TextAreaField('Notes', validators=[Optional()]) submit = SubmitField('Save') class InvestmentTransactionForm(FlaskForm): transaction_type = SelectField('Type', choices=TXN_TYPES, validators=[DataRequired()]) shares = DecimalField('Shares / Units', validators=[DataRequired(), NumberRange(min=0.000001)], places=8) price_per_share = DecimalField('Price per Share', validators=[DataRequired(), NumberRange(min=0)], places=4) fees = DecimalField('Fees', validators=[Optional()], places=2, default=0) date = DateField('Date', validators=[DataRequired()], default=date.today) notes = StringField('Notes', validators=[Optional(), Length(max=255)]) submit = SubmitField('Record Transaction') def _recalc_holding(investment): """ Recalculate shares and avg_cost_basis from transaction log (FIFO approach). """ txns = investment.inv_transactions.order_by( InvestmentTransaction.date.asc() ).all() total_shares = Decimal('0') total_cost = Decimal('0') for txn in txns: shares = Decimal(str(txn.shares)) price = Decimal(str(txn.price_per_share)) fees = Decimal(str(txn.fees or 0)) if txn.transaction_type == 'buy': total_cost += (shares * price) + fees total_shares += shares elif txn.transaction_type == 'sell': if total_shares > 0: avg = total_cost / total_shares total_cost -= avg * min(shares, total_shares) total_shares = max(total_shares - shares, Decimal('0')) elif txn.transaction_type == 'dividend': pass # dividends don't change cost basis elif txn.transaction_type == 'split': # treat split ratio as shares multiplier total_shares += shares investment.shares = total_shares investment.avg_cost_basis = (total_cost / total_shares) if total_shares > 0 else Decimal('0') db.session.commit() @investments_bp.route('/sync-schwab', methods=['POST']) @login_required def sync_schwab(): """Sync balance + positions for every mapped Schwab account, then return here.""" from app.models.schwab_connection import SchwabConnection, SchwabAccount from app.services.schwab_service import sync_account_snapshot import logging log = logging.getLogger(__name__) # Use the active connection for ALL mapped accounts regardless of which # connection_id was stored — accounts may still reference an old connection # after a reconnect if the callback didn't update every record. connection = SchwabConnection.query.filter_by(is_active=True).first() if not connection: flash('No active Schwab connection. Connect at the Schwab page first.', 'warning') return redirect(url_for('investments.index')) accounts = SchwabAccount.query.filter( SchwabAccount.pfm_account_id != None, SchwabAccount.is_active == True, ).all() if not accounts: flash('No Schwab accounts mapped yet. Map them on the Schwab page first.', 'warning') return redirect(url_for('investments.index')) # Patch every account to use the active connection before syncing for sa in accounts: sa.connection = connection ok_accounts, fail_accounts = [], [] total_pos = 0 for sa in accounts: try: _, pos = sync_account_snapshot(sa) total_pos += pos ok_accounts.append(sa.account_name) except Exception as e: log.error('[schwab] sync_schwab failed for %s: %s', sa.account_name, e, exc_info=True) fail_accounts.append(f'{sa.account_name} ({e})') if ok_accounts: flash(f'Synced {len(ok_accounts)} account(s)' f'{f", {total_pos} position(s)" if total_pos else ""}.' f' ({", ".join(ok_accounts)})', 'success') if fail_accounts: flash(f'Failed: {"; ".join(fail_accounts)}', 'danger') if not ok_accounts and not fail_accounts: flash('No mapped Schwab accounts found.', 'info') return redirect(url_for('investments.index')) @investments_bp.route('/') @login_required def index(): portfolio = get_portfolio_summary() # Build allocation chart data chart_labels = [ASSET_TYPE_LABELS.get(a['type'], a['type']) for a in portfolio['allocation']] chart_values = [round(a['value'], 2) for a in portfolio['allocation']] chart_colors = [a['color'] for a in portfolio['allocation']] return render_template('investments/index.html', portfolio=portfolio, chart_labels=chart_labels, chart_values=chart_values, chart_colors=chart_colors, asset_colors=ASSET_COLORS, asset_labels=ASSET_TYPE_LABELS) @investments_bp.route('/new', methods=['GET', 'POST']) @login_required def new(): form = InvestmentForm() if form.validate_on_submit(): inv = Investment( asset_name=form.asset_name.data.strip(), ticker=form.ticker.data.strip().upper() if form.ticker.data else None, asset_type=form.asset_type.data, shares=0, avg_cost_basis=0, notes=form.notes.data, ) db.session.add(inv) db.session.commit() # Fetch initial price if ticker provided if inv.ticker: price = fetch_price(inv.ticker) if price: inv.current_price = price from datetime import datetime inv.last_price_update = datetime.utcnow() db.session.commit() flash(f'"{inv.asset_name}" added. Record your first buy transaction.', 'success') return redirect(url_for('investments.detail', id=inv.id)) return render_template('investments/form.html', form=form, title='Add Investment') @investments_bp.route('/') @login_required def detail(id): inv = db.get_or_404(Investment, id) txns = inv.inv_transactions.order_by( InvestmentTransaction.date.desc() ).all() return render_template('investments/detail.html', inv=inv, txns=txns, asset_labels=ASSET_TYPE_LABELS, asset_colors=ASSET_COLORS) @investments_bp.route('//edit', methods=['GET', 'POST']) @login_required def edit(id): inv = db.get_or_404(Investment, id) form = InvestmentForm(obj=inv) if form.validate_on_submit(): inv.asset_name = form.asset_name.data.strip() inv.ticker = form.ticker.data.strip().upper() if form.ticker.data else None inv.asset_type = form.asset_type.data inv.notes = form.notes.data # Manual override of shares/cost basis — only update if user provided values if form.shares.data is not None: inv.shares = form.shares.data if form.avg_cost_basis.data is not None: inv.avg_cost_basis = form.avg_cost_basis.data db.session.commit() flash('Investment updated.', 'success') return redirect(url_for('investments.detail', id=inv.id)) # Pre-populate shares/cost for edit form if request.method == 'GET': form.shares.data = inv.shares form.avg_cost_basis.data = inv.avg_cost_basis return render_template('investments/form.html', form=form, title='Edit Investment', inv=inv) @investments_bp.route('//delete', methods=['POST']) @login_required def delete(id): inv = db.get_or_404(Investment, id) inv.is_active = False db.session.commit() flash(f'"{inv.asset_name}" removed.', 'info') return redirect(url_for('investments.index')) @investments_bp.route('//transactions/new', methods=['GET', 'POST']) @login_required def add_transaction(id): inv = db.get_or_404(Investment, id) form = InvestmentTransactionForm() if form.validate_on_submit(): shares = form.shares.data price = form.price_per_share.data fees = form.fees.data or 0 total = float(shares) * float(price) txn = InvestmentTransaction( investment_id=inv.id, transaction_type=form.transaction_type.data, shares=shares, price_per_share=price, total_amount=total, fees=fees, date=form.date.data, notes=form.notes.data, ) db.session.add(txn) db.session.commit() # Recalculate holding from transaction log _recalc_holding(inv) # Update current price if no price yet if inv.current_price is None and inv.ticker: p = fetch_price(inv.ticker) if p: from datetime import datetime inv.current_price = p inv.last_price_update = datetime.utcnow() db.session.commit() flash(f'{form.transaction_type.data.title()} transaction recorded.', 'success') return redirect(url_for('investments.detail', id=inv.id)) return render_template('investments/transaction_form.html', form=form, inv=inv, title=f'New Transaction — {inv.asset_name}') @investments_bp.route('/transactions//delete', methods=['POST']) @login_required def delete_transaction(id): txn = db.get_or_404(InvestmentTransaction, id) inv = txn.investment db.session.delete(txn) db.session.commit() _recalc_holding(inv) flash('Transaction deleted.', 'info') return redirect(url_for('investments.detail', id=inv.id)) @investments_bp.route('/refresh-prices', methods=['POST']) @login_required def refresh_prices(): updated = update_prices() if updated: flash(f'Updated prices for: {", ".join(updated.keys())}', 'success') else: flash('No prices updated (no tickers or fetch failed).', 'warning') return redirect(url_for('investments.index')) @investments_bp.route('/api/price/') @login_required def api_price(ticker): """Live price lookup for a ticker — used in the add/edit investment form.""" ticker = ticker.upper().strip() error = None price = None try: price = fetch_price(ticker) if price is None: error = f'No data returned for {ticker}. Check the ticker format.' except Exception as e: error = str(e) return jsonify({ 'ticker': ticker, 'price': price, 'error': error, }) @investments_bp.route('/api/daychange/') @login_required def api_day_change(ticker): """ Lightweight endpoint: return today's open-to-current day change only. Used by the portfolio page to load change badges quickly. """ data = fetch_day_change(ticker.upper().strip()) if data is None: return jsonify({'error': f'No data for {ticker}'}), 404 return jsonify(data) @investments_bp.route('/api/price-alerts') @login_required def api_price_alerts(): """Return today's price-alert list (holdings that moved >= 5% intraday).""" return jsonify({'alerts': get_price_alerts()}) @investments_bp.route('/api/history/') @login_required def api_price_history(ticker): """ Return OHLC history + day/period change for a ticker. Query param: tf = 1W | 1M | 3M | 6M | 1Y (default 1M) Used by the portfolio page inline charts. """ tf = request.args.get('tf', '1M').upper() if tf not in ('1W', '1M', '3M', '6M', '1Y'): tf = '1M' data = fetch_price_history(ticker.upper().strip(), tf) if data is None: return jsonify({'error': f'No history data available for {ticker}'}), 404 return jsonify(data)