Files

180 lines
6.2 KiB
Python

import logging
from django.shortcuts import get_object_or_404
from rest_framework import status, viewsets
from rest_framework.decorators import action
from rest_framework.response import Response
from rest_framework.views import APIView
from .models import BenchmarkPrice, CashFlow, Portfolio, PortfolioSnapshot, Stock, Transaction
from .serializers import (
AIUpdateSerializer,
BenchmarkPriceSerializer,
CashFlowSerializer,
PortfolioListSerializer,
PortfolioSerializer,
PortfolioSnapshotSerializer,
StockSerializer,
TransactionSerializer,
)
from .services import (
ai_update_holdings,
get_agent_summary,
get_cashflow_adjusted_performance,
get_portfolio_value,
get_risk_summary,
)
logger = logging.getLogger(__name__)
class PortfolioViewSet(viewsets.ModelViewSet):
queryset = Portfolio.objects.prefetch_related('stocks').all()
def get_serializer_class(self):
if self.action == 'list':
return PortfolioListSerializer
return PortfolioSerializer
@action(detail=True, methods=['get'], url_path='holdings')
def holdings(self, request, pk=None):
"""Return holdings with real-time prices."""
portfolio = self.get_object()
reference_date = request.query_params.get('reference_date')
try:
data = get_portfolio_value(portfolio, reference_date=reference_date)
return Response(data)
except Exception as exc:
logger.error("get_portfolio_value failed for %s: %s", portfolio.id, exc, exc_info=True)
return Response({'error': str(exc)}, status=status.HTTP_500_INTERNAL_SERVER_ERROR)
@action(detail=True, methods=['get'], url_path='transactions')
def transactions(self, request, pk=None):
"""List all transactions for this portfolio."""
portfolio = self.get_object()
txs = portfolio.transactions.all().order_by('-date', '-created_at')
serializer = TransactionSerializer(txs, many=True)
return Response(serializer.data)
class StockViewSet(viewsets.ModelViewSet):
queryset = Stock.objects.select_related('portfolio').all()
serializer_class = StockSerializer
def get_queryset(self):
qs = super().get_queryset()
portfolio_id = self.request.query_params.get('portfolio')
if portfolio_id:
qs = qs.filter(portfolio_id=portfolio_id)
return qs
class TransactionViewSet(viewsets.ModelViewSet):
queryset = Transaction.objects.select_related('portfolio').all()
serializer_class = TransactionSerializer
def get_queryset(self):
qs = super().get_queryset()
portfolio_id = self.request.query_params.get('portfolio')
if portfolio_id:
qs = qs.filter(portfolio_id=portfolio_id)
stock_code = self.request.query_params.get('stock_code')
if stock_code:
qs = qs.filter(stock_code=stock_code.upper())
return qs.order_by('-date', '-created_at')
def perform_create(self, serializer):
serializer.save(stock_code=serializer.validated_data['stock_code'].upper())
class CashFlowViewSet(viewsets.ModelViewSet):
queryset = CashFlow.objects.select_related('portfolio').all()
serializer_class = CashFlowSerializer
def get_queryset(self):
qs = super().get_queryset()
portfolio_id = self.request.query_params.get('portfolio')
if portfolio_id:
qs = qs.filter(portfolio_id=portfolio_id)
flow_type = self.request.query_params.get('flow_type')
if flow_type:
qs = qs.filter(flow_type=flow_type.upper())
return qs.order_by('-date', '-created_at')
class PortfolioSnapshotViewSet(viewsets.ReadOnlyModelViewSet):
queryset = PortfolioSnapshot.objects.select_related('portfolio').all()
serializer_class = PortfolioSnapshotSerializer
def get_queryset(self):
qs = super().get_queryset()
portfolio_id = self.request.query_params.get('portfolio')
if portfolio_id:
qs = qs.filter(portfolio_id=portfolio_id)
return qs.order_by('-captured_at')
class BenchmarkPriceViewSet(viewsets.ReadOnlyModelViewSet):
queryset = BenchmarkPrice.objects.all()
serializer_class = BenchmarkPriceSerializer
def get_queryset(self):
qs = super().get_queryset()
ticker = self.request.query_params.get('ticker')
if ticker:
qs = qs.filter(ticker=ticker.upper())
return qs.order_by('ticker', 'date')
class AIUpdateView(APIView):
"""POST /api/invest/ai-update/ — Sync portfolio holdings (quantity only required)."""
def post(self, request):
serializer = AIUpdateSerializer(data=request.data)
if not serializer.is_valid():
return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST)
data = serializer.validated_data
portfolio = get_object_or_404(Portfolio, pk=data['portfolio_id'])
try:
result = ai_update_holdings(
portfolio=portfolio,
holdings=data['holdings'],
reset=data['reset'],
)
except Exception as exc:
logger.error("ai_update_holdings failed: %s", exc, exc_info=True)
return Response({'error': str(exc)}, status=status.HTTP_400_BAD_REQUEST)
return Response(result, status=status.HTTP_200_OK)
class AgentSummaryView(APIView):
"""GET /api/invest/agent/summary/ — agent-friendly portfolio summary."""
def get(self, request):
return Response(get_agent_summary())
class PerformanceView(APIView):
"""GET /api/invest/performance/?start=YYYY-MM-DD&end=YYYY-MM-DD&benchmarks=QQQ,SPY"""
def get(self, request):
benchmarks = request.query_params.get('benchmarks', 'QQQ,SPY')
tickers = [item.strip().upper() for item in benchmarks.split(',') if item.strip()]
return Response(
get_cashflow_adjusted_performance(
start=request.query_params.get('start'),
end=request.query_params.get('end'),
benchmark_tickers=tickers,
)
)
class RiskView(APIView):
"""GET /api/invest/risk/ — concentration and theme exposure."""
def get(self, request):
return Response(get_risk_summary())